Compare commits

...
Author SHA1 Message Date
yyhandGitHub b95d313128 refactor(web): improve workflow block selector (#39377) 2026-07-21 14:18:38 +00:00
github-actions[bot]GitHubclaude[bot] <41898282+claude[bot]@users.noreply.github.com>yyh
45ac70132c chore(i18n): sync translations with en-US (#39362)
Co-authored-by: claude[bot] <41898282+claude[bot]@users.noreply.github.com>
Co-authored-by: yyh <92089059+lyzno1@users.noreply.github.com>
2026-07-21 11:54:20 +00:00
JoelandGitHub 6e772259bb fix: update the agent roster documentation link (#39368) 2026-07-21 09:56:21 +00:00
QuantumGhostandGitHub af465fb527 fix(api): scope HITL snapshot message lookup to utilize db index (#39351) 2026-07-21 09:47:17 +00:00
Xiyuan ChenGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
62bbc0dbeb feat(webapp-access): add WEBAPP_PUBLIC_ACCESS_ENABLED toggle (#38567)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-21 08:46:28 +00:00
zyssyz123andGitHub ed1289281d fix(agent): include workflow runs in monitoring stats (#39354) 2026-07-21 08:33:23 +00:00
Asuka MinatoandGitHub e37f32e4a0 test: add SQLite session fixture for providers (#38982) 2026-07-21 08:29:44 +00:00
Asuka MinatoandGitHub 150613cc05 test: use sqlite3 session in test_build_from_mapping (#38769) 2026-07-21 08:21:36 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
a5194ff17b test: move human input delivery coverage to unit tests (#38926)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-21 08:18:32 +00:00
Asuka MinatoandGitHub cdaf5d4562 test: use sqlite3 session in test_tool_file_manager (#38771) 2026-07-21 08:10:23 +00:00
Asuka MinatoandGitHub 145b0e622e test: use sqlite3 session in test_workspace_account (#38757) 2026-07-21 08:07:57 +00:00
Asuka MinatoandGitHub 4bab62d458 test: use sqlite3 session in test_agent_app_parameters (#38759) 2026-07-21 07:59:00 +00:00
Asuka MinatoandGitHub b99f8941bb test: use sqlite3 session in test_database_retrieval (#38743) 2026-07-21 07:56:34 +00:00
Asuka MinatoandGitHub 32d2594877 test: move app import API coverage to unit tests (#38920) 2026-07-21 07:54:32 +00:00
Asuka MinatoandGitHub 7b188287da test: use sqlite3 session in test_accounts (#38752) 2026-07-21 07:50:05 +00:00
Asuka MinatoandGitHub f6db23444d test: move API token service coverage to unit tests (#38918) 2026-07-21 07:48:10 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
a5c989d82b test: use sqlite3 session in test_conversation_service (#38774)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-21 07:38:53 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2622b1ce0a test: use sqlite3 session in test_based_generate_task_pipeline (#38738)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-21 07:37:20 +00:00
Asuka MinatoandGitHub 1e672530c4 test: use sqlite3 session in test_built_in_retrieval (#38735) 2026-07-21 07:36:42 +00:00
Asuka MinatoandGitHub 852169d99c test: use SQLite sessions in unit misc (#39118) 2026-07-21 07:27:13 +00:00
Asuka MinatoandGitHub d7860571aa test: use SQLite sessions in services workflow (#39115) 2026-07-21 07:24:52 +00:00
Asuka MinatoandGitHub 3683831738 test: use sqlite3 session in test_account_service (#38751) 2026-07-21 07:22:05 +00:00
非法操作andGitHub f5024f5440 fix(workflow): avoid false draft save error during Strict Mode initia… (#39353) 2026-07-21 07:11:19 +00:00
JoelGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
4f8477427e fix: not support jump to detail in marketplace (#39347)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-21 07:01:00 +00:00
非法操作andGitHub 8d406aec31 fix(dev): use gevent WebSocket server for local API (#39349) 2026-07-21 06:42:55 +00:00
落尘andGitHub f8754286ec fix(watercrawl): handle non-json error responses (#37514) 2026-07-21 06:41:06 +00:00
yyhandGitHub 7da6b2ce36 refactor(web): remove permission catalog service hooks (#39342) 2026-07-21 06:01:00 +00:00
aa718add24 fix(api): return 400 instead of 500 when 'file' is missing from trial-app audio-to-text (#39305)
Co-authored-by: Harsh Kashyap <Harsh23Kashyap@users.noreply.github.com>
2026-07-21 05:59:48 +00:00
yyhandGitHub fcc1d39f5e fix(web): restore billing access boundary (#39343) 2026-07-21 05:35:36 +00:00
yyhandGitHub 861831f176 fix(api): keep workflow log archives owner-admin only (#39338) 2026-07-21 04:31:14 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
d333aa31ea test: use SQLite sessions in services workflow (#39116)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-21 03:40:02 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
97e6a14c13 test: use SQLite sessions in unit misc (#39119)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-21 03:20:59 +00:00
shaktiandGitHub 6a338654fb fix: return 400 instead of 500 when file field is missing in audio-to-text endpoints (#39303) (#39322) 2026-07-21 02:19:24 +00:00
ChecoandGitHub 8f075c82bc fix(configs): reject host:port-shaped PLUGIN_REMOTE_INSTALL_PORT with actionable hint (#39329)
Signed-off-by: sergioperezcheco <checo520@outlook.com>
2026-07-21 02:18:29 +00:00
林玮 (Jade Lin)andGitHub 28d17da3b1 fix(api): use resource tenant for draft variable files (#39307) 2026-07-21 02:07:47 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
891b2dc537 test: use SQLite sessions in unit misc (#39120)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-21 01:51:24 +00:00
WH-2099andGitHub 75f7069541 ci: simplify test workflow execution (#39318) 2026-07-21 00:18:59 +00:00
yyhandGitHub b7193d1cba refactor(web): remove basePath from hydration boundary (#39308) 2026-07-20 14:46:37 +00:00
wangxiaoleiGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
ef0115d340 chore: when dataset permission is not all team, update rbac config to… (#39227)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-20 12:59:29 +00:00
3afb9b3230 fix(web): stop doubling basePath in auth refresh redirects (#39273)
Co-authored-by: yyh <yuanyouhuilyz@gmail.com>
2026-07-20 12:22:40 +00:00
JoelandGitHub 0862641533 feat: support export agent dsl in sidebar (#39299) 2026-07-20 10:41:26 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
cbd7872520 test: use SQLite sessions in unit misc (#39121)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-20 10:38:52 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
23882a704e test: use SQLite sessions in unit misc (#39122)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-20 10:38:49 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
1bd5254ea7 test: use SQLite sessions in unit misc (#39123)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-20 10:38:46 +00:00
yyhandGitHub e13069dbe3 revert(ci): restore default Depot runners (#39297) 2026-07-20 10:18:52 +00:00
github-actions[bot]GitHubclaude[bot] <41898282+claude[bot]@users.noreply.github.com>yyh
edd3d4a218 chore(i18n): sync translations with en-US (#39296)
Co-authored-by: claude[bot] <41898282+claude[bot]@users.noreply.github.com>
Co-authored-by: yyh <92089059+lyzno1@users.noreply.github.com>
2026-07-20 10:15:37 +00:00
JoelandGitHub 79751f7622 chore: display non-LLM settings in integration tool details (#39295) 2026-07-20 09:52:54 +00:00
JingyiGitHubhjlarryautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>yyhyyh
cb64446fa3 feat(web): add step-by-step tour shell (#38785)
Co-authored-by: hjlarry <hjlarry@163.com>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: yyh <yuanyouhuilyz@gmail.com>
Co-authored-by: yyh <92089059+lyzno1@users.noreply.github.com>
2026-07-20 09:23:25 +00:00
JoelandGitHub e77ee82526 fix: show persistent error when workflow draft save fails (#39293) 2026-07-20 08:48:38 +00:00
yyhandGitHub b3f163fb0f chore(deps): update pnpm and workspace dependencies (#39292) 2026-07-20 08:43:21 +00:00
yyhandGitHub 873ef2b592 docs: refine frontend testing guidance (#39290) 2026-07-20 08:18:26 +00:00
Stephen ZhouGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
43d725a1ce feat(dataset): proxy KnowledgeFS Console requests (#39158)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-20 08:10:06 +00:00
JoelandGitHub c8ccfba960 fix: the input of agent chat would be replaced if conversation history updated (#39289) 2026-07-20 08:07:12 +00:00
eee8d6cf7b feat(web): add webapp access control to agent access point (#39285)
Co-authored-by: yyh <yuanyouhuilyz@gmail.com>
2026-07-20 08:05:29 +00:00
zyssyz123andGitHub 11016563b1 fix(agent): preserve complete model usage pricing (#39201) 2026-07-20 08:03:13 +00:00
Ingram ZandGitHub b7edc5ef0c refactor: use catalog for workflow archive maintenance (#39001) 2026-07-20 07:59:29 +00:00
林玮 (Jade Lin)GitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
f93a5b95c6 fix(api, web): scope trial app file uploads to the app tenant (#39149)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-20 07:00:45 +00:00
Kevin9703GitHubhjlarryautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>yyhyyh
6022939bf0 fix: prevent hidden-tab collaboration leader from saving stale drafts (#38997)
Co-authored-by: hjlarry <hjlarry@163.com>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: yyh <yuanyouhuilyz@gmail.com>
Co-authored-by: yyh <92089059+lyzno1@users.noreply.github.com>
2026-07-20 06:58:01 +00:00
非法操作andGitHub 30df1433cc fix: can't create chat app when get tenant default model schema error (#39266) 2026-07-20 06:47:16 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
aa7b65c602 chore: bump the github-actions-dependencies group with 5 updates (#39247)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-07-20 06:47:01 +00:00
yyhandGitHub ba24745c07 fix(web): restore permission selector semantics (#39269)
Signed-off-by: yyh <yuanyouhuilyz@gmail.com>
2026-07-20 06:30:09 +00:00
38ea0de907 refactor(web): avoid direct setState in effect in jina reader crawler (#25194) (#39216)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
Co-authored-by: Wu Tianwei <30284043+WTW0313@users.noreply.github.com>
2026-07-20 06:28:08 +00:00
JoelandGitHub bbc31a3403 fix: text generation web app keeps loading after timeout (#39268) 2026-07-20 06:27:32 +00:00
605c0fe02b refactor(web): avoid setState in website crawler effects (#39237)
Co-authored-by: Wu Tianwei <30284043+WTW0313@users.noreply.github.com>
2026-07-20 06:26:50 +00:00
yyhandGitHub 15674a9db6 test(web): align test state boundaries (#39239) 2026-07-20 06:06:17 +00:00
非法操作GitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
9a2d039e93 chore: the node should not draggable when preview a workflow (#39131)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-20 05:59:14 +00:00
yyhandGitHub 58d1dd8873 ci: route default Depot jobs to 4-core runners (#39277) 2026-07-20 05:42:56 +00:00
ChecoandGitHub 5ea884f799 fix(dataset_service): use flush() instead of commit() to preserve caller transaction (#39223) 2026-07-19 07:00:50 +00:00
yyhGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2ec34b2cfb feat(workflow): improve block selector navigation and previews (#39212)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-18 09:56:36 +00:00
fa65f45bff refactor(web): align integrations sidebar components (#39202)
Co-authored-by: Jingyi <jingyi.qi@dify.ai>
2026-07-17 22:52:47 +00:00
yyhGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
63ca2b94b5 refactor(web): use generated OAuth query options (#39204)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-17 14:50:56 +00:00
yyhandGitHub a7ef41a8f5 refactor(web): use generated dataset card queries (#39203) 2026-07-17 14:50:40 +00:00
Yunlu WenandGitHub 5c6372d2f7 chore: bump version to 1.16.0 (#39196) 2026-07-17 09:28:32 +00:00
yyhandGitHub 5fd06fafe0 test(web): remove low-value unit tests (#39193) 2026-07-17 08:59:26 +00:00
zyssyz123andGitHub af4b65b295 fix(agent): allow dangling refs in workflow drafts (#39192) 2026-07-17 08:42:58 +00:00
Yunlu WenGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
6e5fc1081b feat: refact agent shell output (#39190)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-17 08:27:00 +00:00
JoelandGitHub 0e84ae7338 fix: can not typing in the change node search field (#39189) 2026-07-17 08:01:46 +00:00
zyssyz123andGitHub 5a792945f5 fix(agent): bound streams and cancel zombie workflow runs (#39186) 2026-07-17 07:46:42 +00:00
Wu TianweiGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2f3c785f27 feat: add Responses API tip for GPT-5.6 and later models in multiple languages (#39188)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-17 07:34:38 +00:00
JoelandGitHub 38045078dc fix: can not create app from learn dify in app list page (#39187) 2026-07-17 07:20:31 +00:00
sahasweetyGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
83f40b85f4 fix: remove patch logger in test, use caplog instead (Closes #37468) (#39159)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-17 06:22:20 +00:00
非法操作andGitHub aff8a49bbf fix: change email not work (#39182) 2026-07-17 06:05:20 +00:00
JoelandGitHub b0329cd53c chore: add agent tooltip (#39183) 2026-07-17 06:00:29 +00:00
yyhandGitHub 2a008423d8 chore: upgrade Vite+ to 0.2.5 (#39181) 2026-07-17 05:39:02 +00:00
yyhandGitHub f962c9e47a fix(web): settle running nodes after workflow failure (#39173) 2026-07-17 05:15:01 +00:00
yyhandGitHub af7e59de7c fix(web): show current profile in main navigation (#39172) 2026-07-17 04:11:35 +00:00
非法操作andGitHub 61b07ab17d feat: default plugins for new user (#39167) 2026-07-17 02:49:58 +00:00
JoelandGitHub a7aff83d52 chore: remove useless agent tip (#39171) 2026-07-17 02:47:57 +00:00
9de7e0fe44 perf(workflow-generator): parallelize node config generation (#38975)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
2026-07-17 02:34:15 +00:00
1921 changed files with 46079 additions and 121802 deletions
+2 -3
View File
@@ -9,12 +9,11 @@ Use this skill for Vitest work under `web/` and `packages/dify-ui/`. Do not use
## Required Source
Before writing, changing, or reviewing frontend tests, read `web/docs/test.md` completely. It is the single source of truth. This skill defines the execution workflow and must not add requirements that conflict with or duplicate that guide.
Before writing, changing, or reviewing frontend tests, read `web/docs/test.md` completely. It is the single source of truth. This skill provides an execution checklist and must not redefine or extend that policy.
## Workflow
1. Read the source, its behavior owner, nearby specs, and relevant public dependencies.
1. Identify whether the contract belongs in `web/`, Dify UI Browser Mode, or a styled Storybook test.
1. Apply the canonical guide to decide whether a test is needed and choose its boundary.
1. For a behavior change or bug fix, write or identify the failing scenario first when practical.
1. Implement one coherent scenario at a time and run the focused spec before expanding scope.
@@ -33,4 +32,4 @@ vp test run path/to/spec-or-directory
vp test run --project unit src/path/to/spec
```
For styled Dify UI behavior, run `vp test --project storybook --run`. Run broader checks only after the focused behavior passes.
Run Dify UI Storybook tests with `vp test --project storybook --run`. Run broader checks only after the focused behavior passes.
-3
View File
@@ -47,9 +47,6 @@ jobs:
- name: Install dependencies
run: uv sync --project api --dev
- name: Run dify config tests
run: uv run --project api pytest api/tests/unit_tests/configs/test_env_consistency.py
- name: Run Unit Tests
run: |
uv run --project api pytest \
+1 -2
View File
@@ -27,7 +27,7 @@ jobs:
steps:
- id: skip_check
continue-on-error: true
uses: fkirc/skip-duplicate-actions@f75f66ce1886f00957d99748a42c724f4330bdcf # v5.3.1
uses: fkirc/skip-duplicate-actions@b974a9395958c231af965b70070979a577efa578 # v5.3.2
with:
cancel_others: 'true'
concurrent_skipping: same_content_newer
@@ -81,7 +81,6 @@ jobs:
- '.npmrc'
- '.nvmrc'
- '.github/workflows/cli-tests.yml'
- '.github/workflows/cli-docker-build.yml'
- '.github/actions/setup-web/**'
web:
- 'web/**'
+4 -4
View File
@@ -27,7 +27,7 @@ jobs:
persist-credentials: false
- name: Setup Go
uses: actions/setup-go@d35c59abb061a4a6fb18e82ac0862c26744d6ab5 # v5.5.0
uses: actions/setup-go@b7ad1dad31e06c5925ef5d2fc7ad053ef454303e # v7.0.0
with:
go-version-file: dify-agent-runtime/go.mod
cache-dependency-path: dify-agent-runtime/go.sum
@@ -51,13 +51,13 @@ jobs:
persist-credentials: false
- name: Setup Go
uses: actions/setup-go@d35c59abb061a4a6fb18e82ac0862c26744d6ab5 # v5.5.0
uses: actions/setup-go@b7ad1dad31e06c5925ef5d2fc7ad053ef454303e # v7.0.0
with:
go-version-file: dify-agent-runtime/go.mod
cache-dependency-path: dify-agent-runtime/go.sum
- name: Run golangci-lint
uses: golangci/golangci-lint-action@4afd733a84b1f43292c63897423277bb7f4313a9 # v6.5.0
uses: golangci/golangci-lint-action@ba0d7d2ec06a0ea1cb5fa41b2e4a3ab91d21278a # v6.5.0
with:
working-directory: dify-agent-runtime
version: latest
@@ -78,7 +78,7 @@ jobs:
persist-credentials: false
- name: Setup Go
uses: actions/setup-go@d35c59abb061a4a6fb18e82ac0862c26744d6ab5 # v5.5.0
uses: actions/setup-go@b7ad1dad31e06c5925ef5d2fc7ad053ef454303e # v7.0.0
with:
go-version-file: dify-agent-runtime/go.mod
cache-dependency-path: dify-agent-runtime/go.sum
+1 -1
View File
@@ -29,7 +29,7 @@ jobs:
persist-credentials: false
- name: Use Node.js
uses: actions/setup-node@48b55a011bda9f5d6aeb4c2d9c7362e8dae4041e # v6.4.0
uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7.0.0
with:
node-version: 22
cache: ''
+1 -1
View File
@@ -158,7 +158,7 @@ jobs:
- name: Run Claude Code for Translation Sync
if: steps.context.outputs.CHANGED_FILES != ''
uses: anthropics/claude-code-action@e90deca47693f9457b72f2b53c17d7c445a87342 # v1.0.171
uses: anthropics/claude-code-action@af0559ee4f514d1ef21826982bed13f7edc3c35e # v1.0.178
with:
anthropic_api_key: ${{ secrets.ANTHROPIC_API_KEY }}
github_token: ${{ secrets.GITHUB_TOKEN }}
-10
View File
@@ -135,16 +135,6 @@ jobs:
exit 1
fi
if [[ -d cucumber-report ]]; then
rm -rf cucumber-report-non-external
mv cucumber-report cucumber-report-non-external
fi
if [[ -d .logs ]]; then
rm -rf .logs-non-external
mv .logs .logs-non-external
fi
teardown_external_runtime() {
local run_status=$?
trap - EXIT
-2
View File
@@ -17,8 +17,6 @@ jobs:
test:
name: Web Tests (${{ matrix.shardIndex }}/${{ matrix.shardTotal }})
runs-on: depot-ubuntu-24.04-4
env:
VITEST_COVERAGE_SCOPE: app-components
strategy:
fail-fast: false
matrix:
+4 -18
View File
@@ -5,7 +5,9 @@
"name": "Python: API (gevent)",
"type": "debugpy",
"request": "launch",
"program": "${workspaceFolder}/api/app.py",
"module": "gevent.monkey",
"args": ["--module", "app"],
"gevent": true,
"jinja": true,
"justMyCode": true,
"cwd": "${workspaceFolder}/api",
@@ -33,22 +35,6 @@
"justMyCode": false,
"cwd": "${workspaceFolder}/api",
"python": "${workspaceFolder}/api/.venv/bin/python"
},
{
"name": "Next.js: debug full stack",
"type": "node",
"request": "launch",
"program": "${workspaceFolder}/web/node_modules/next/dist/bin/next",
"runtimeArgs": ["--inspect"],
"skipFiles": ["<node_internals>/**"],
"serverReadyAction": {
"action": "debugWithChrome",
"killOnServerStop": true,
"pattern": "- Local:.+(https?://.+)",
"uriFormat": "%s",
"webRoot": "${workspaceFolder}/web"
},
"cwd": "${workspaceFolder}/web"
}
}
]
}
+2
View File
@@ -107,6 +107,7 @@ test:
echo "Target: $(TARGET_TESTS)"; \
uv run --project api --dev pytest $(TARGET_TESTS); \
else \
set -e; \
echo "Running backend unit tests"; \
uv run --project api --dev pytest -p no:benchmark --timeout "$${PYTEST_TIMEOUT:-20}" -n auto \
api/tests/unit_tests \
@@ -124,6 +125,7 @@ test-all:
echo "Target: $(TARGET_TESTS)"; \
uv run --project api --dev pytest $(TARGET_TESTS); \
else \
set -e; \
echo "Running backend unit tests"; \
uv run --project api --dev pytest -p no:benchmark --timeout "$${PYTEST_TIMEOUT:-20}" -n auto \
api/tests/unit_tests \
+19
View File
@@ -561,6 +561,8 @@ WORKFLOW_MAX_EXECUTION_STEPS=500
WORKFLOW_MAX_EXECUTION_TIME=1200
WORKFLOW_CALL_MAX_DEPTH=5
MAX_VARIABLE_SIZE=204800
# Maximum concurrent node-builder LLM calls per workflow generation request
WORKFLOW_GENERATOR_NODE_BUILDER_MAX_WORKERS=6
# GraphEngine Worker Pool Configuration
# Minimum number of workers per GraphEngine instance (default: 1)
@@ -665,10 +667,27 @@ PLUGIN_REMOTE_INSTALL_HOST=localhost
PLUGIN_MAX_PACKAGE_SIZE=15728640
PLUGIN_MODEL_SCHEMA_CACHE_TTL=3600
PLUGIN_MODEL_PROVIDERS_CACHE_TTL=86400
# Comma-separated marketplace plugin IDs whose latest versions are installed for newly registered users.
# Example: langgenius/openai,langgenius/gemini
NEW_USER_DEFAULT_PLUGIN_IDS=
# Comma-separated model_type:provider:model entries assigned after default plugins finish installing.
# Example: llm:langgenius/openai/openai:gpt-4o-mini,text-embedding:langgenius/openai/openai:text-embedding-3-small
NEW_USER_DEFAULT_MODELS=
INNER_API_KEY_FOR_PLUGIN=QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1
# Dify Agent backend
AGENT_BACKEND_BASE_URL=http://localhost:5050
AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS=30
AGENT_BACKEND_STREAM_MAX_RECONNECTS=3
AGENT_BACKEND_RUN_TIMEOUT_SECONDS=1200
# KnowledgeFS (Dataset 2.0)
KNOWLEDGE_FS_ENABLED=false
KNOWLEDGE_FS_BASE_URL=
# Shared with KnowledgeFS; use at least 32 random characters.
KNOWLEDGE_FS_JWT_SECRET=
KNOWLEDGE_FS_SSE_READ_TIMEOUT_SECONDS=300
KNOWLEDGE_FS_TIMEOUT_SECONDS=10
# Marketplace configuration
MARKETPLACE_ENABLED=true
+19
View File
@@ -1,5 +1,24 @@
from __future__ import annotations
# ``python -m app`` (docker DEBUG=true, or IDE debugging) serves through the
# gevent pywsgi server at the bottom of this file, so the stdlib must be
# monkey-patched BEFORE any other import pulls in sockets or locks. Without
# this, every request runs as a greenlet on one OS thread while blocking
# calls (LLM invokes, ``Future.result`` waits, DB I/O) pin that thread — the
# whole process freezes until the call returns. Gunicorn and Celery apply
# their own patching (see gunicorn.conf.py / celery_entrypoint.py), and
# ``flask run`` uses real Werkzeug threads, so both skip this branch.
if __name__ == "__main__":
from gevent import monkey
monkey.patch_all()
import psycogreen.gevent as psycogreen_gevent
from grpc.experimental import gevent as grpc_gevent
grpc_gevent.init_gevent()
psycogreen_gevent.patch_psycopg()
import logging
import sys
from typing import TYPE_CHECKING, cast
+40 -6
View File
@@ -8,7 +8,7 @@ creating another wire contract.
from __future__ import annotations
from collections.abc import Iterator
from collections.abc import Callable, Iterator
from typing import Protocol
from dify_agent.client import (
@@ -45,7 +45,13 @@ class AgentBackendRunClient(Protocol):
def cancel_run(self, run_id: str, request: CancelRunRequest | None = None) -> CancelRunResponse:
"""Request explicit cancellation for one Agent backend run."""
def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]:
def stream_events(
self,
run_id: str,
*,
after: str | None = None,
should_stop: Callable[[], bool] | None = None,
) -> Iterator[RunEvent]:
"""Yield public ``dify-agent`` run events in stream order."""
def wait_run(self, run_id: str, *, timeout_seconds: float | None = None) -> RunStatusResponse:
@@ -61,7 +67,15 @@ class _DifyAgentSyncClient(Protocol):
def cancel_run_sync(self, run_id: str, request: CancelRunRequest | None = None) -> CancelRunResponse:
"""Cancel one run synchronously."""
def stream_events_sync(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]:
def stream_events_sync(
self,
run_id: str,
*,
after: str | None = None,
max_reconnects: int | None = None,
timeout_seconds: float | None = None,
should_stop: Callable[[], bool] | None = None,
) -> Iterator[RunEvent]:
"""Stream run events synchronously."""
def wait_run_sync(self, run_id: str, *, timeout_seconds: float | None = None) -> RunStatusResponse:
@@ -73,8 +87,16 @@ class DifyAgentBackendRunClient:
client: _DifyAgentSyncClient
def __init__(self, client: _DifyAgentSyncClient) -> None:
def __init__(
self,
client: _DifyAgentSyncClient,
*,
stream_max_reconnects: int = 3,
stream_timeout_seconds: float = 1200,
) -> None:
self.client = client
self._stream_max_reconnects = stream_max_reconnects
self._stream_timeout_seconds = stream_timeout_seconds
def create_run(self, request: CreateRunRequest) -> CreateRunResponse:
"""Create one run through ``POST /runs`` and normalize client exceptions."""
@@ -90,10 +112,22 @@ class DifyAgentBackendRunClient:
except Exception as exc:
raise _normalize_dify_agent_error(exc) from exc
def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]:
def stream_events(
self,
run_id: str,
*,
after: str | None = None,
should_stop: Callable[[], bool] | None = None,
) -> Iterator[RunEvent]:
"""Stream run events from ``/events/sse`` with the wrapped client's reconnect policy."""
try:
yield from self.client.stream_events_sync(run_id, after=after)
yield from self.client.stream_events_sync(
run_id,
after=after,
max_reconnects=self._stream_max_reconnects,
timeout_seconds=self._stream_timeout_seconds,
should_stop=should_stop,
)
except Exception as exc:
raise _normalize_dify_agent_error(exc) from exc
+8 -1
View File
@@ -13,10 +13,17 @@ def create_agent_backend_run_client(
base_url: str | None = None,
use_fake: bool = False,
fake_scenario: str | FakeAgentBackendScenario = FakeAgentBackendScenario.SUCCESS,
stream_read_timeout_seconds: float = 30,
stream_max_reconnects: int = 3,
stream_run_timeout_seconds: float = 1200,
) -> AgentBackendRunClient:
"""Create the API-side run client without hiding the ``dify-agent`` protocol."""
if use_fake:
return FakeAgentBackendRunClient(scenario=FakeAgentBackendScenario(fake_scenario))
if base_url is None:
raise ValueError("base_url is required when creating a real Agent backend client")
return DifyAgentBackendRunClient(Client(base_url=base_url))
return DifyAgentBackendRunClient(
Client(base_url=base_url, stream_timeout=stream_read_timeout_seconds),
stream_max_reconnects=stream_max_reconnects,
stream_timeout_seconds=stream_run_timeout_seconds,
)
+10 -2
View File
@@ -7,7 +7,7 @@ separate ``agent-backend.v1`` event stream.
from __future__ import annotations
from collections.abc import Iterator
from collections.abc import Callable, Iterator
from datetime import UTC, datetime
from enum import StrEnum
@@ -69,9 +69,17 @@ class FakeAgentBackendRunClient:
del request
return CancelRunResponse(run_id=run_id, status="cancelled")
def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]:
def stream_events(
self,
run_id: str,
*,
after: str | None = None,
should_stop: Callable[[], bool] | None = None,
) -> Iterator[RunEvent]:
"""Yield the deterministic public ``RunEvent`` sequence for ``run_id``."""
for event in self._events(run_id):
if should_stop is not None and should_stop():
return
if after is not None and event.id is not None and event.id <= after:
continue
yield event
+209 -63
View File
@@ -1,6 +1,8 @@
import datetime
import logging
import re
import time
import uuid
from collections.abc import Callable
from typing import TypedDict
@@ -21,6 +23,7 @@ from tasks.remove_app_and_related_data_task import delete_draft_variables_batch
logger = logging.getLogger(__name__)
_HEX_PREFIXES = tuple("0123456789abcdef")
_TARGET_MONTH_PATTERN = re.compile(r"^\d{4}-(0[1-9]|1[0-2])$")
class WorkflowRunArchivePlanRow(TypedDict):
@@ -66,6 +69,7 @@ def _parse_tenant_prefixes(prefixes: str | None) -> list[str]:
def _parse_comma_separated_ids(raw_ids: str | None, *, param_name: str) -> list[str] | None:
"""Keep an omitted scope unset while rejecting an explicitly empty scope."""
if raw_ids is None:
return None
parsed = sorted({raw_id.strip() for raw_id in raw_ids.split(",") if raw_id.strip()})
@@ -74,6 +78,27 @@ def _parse_comma_separated_ids(raw_ids: str | None, *, param_name: str) -> list[
return parsed
def _parse_archive_target_month(target_month: str) -> tuple[int, int]:
"""Validate the V2 catalog month selector and return its numeric components."""
if not _TARGET_MONTH_PATTERN.fullmatch(target_month):
raise click.BadParameter("target-month must use YYYY-MM format", param_hint="--target-month")
year_text, month_text = target_month.split("-", maxsplit=1)
return int(year_text), int(month_text)
def _parse_archive_catalog_cursor(after_catalog_id: str | None) -> str | None:
"""Normalize the exclusive V2 catalog keyset cursor when one is provided."""
if after_catalog_id is None:
return None
try:
return str(uuid.UUID(after_catalog_id))
except ValueError as exc:
raise click.BadParameter(
"after-catalog-id must be a UUID returned by the same V2 operation and scope",
param_hint="--after-catalog-id",
) from exc
def _get_archive_candidate_tenant_ids_by_prefix(
session: Session,
prefix: str,
@@ -810,9 +835,11 @@ def backfill_workflow_run_archive_bundles(
click.echo(click.style(f" ... and {len(summary.errors) - 10} more failures", fg="red"))
def _echo_bundle_archive_operation_summary(summary) -> None:
def _echo_bundle_archive_operation_summary(summary, *, dry_run: bool) -> None:
status = "completed successfully" if summary.bundles_failed == 0 else "completed with failures"
fg = "green" if summary.bundles_failed == 0 else "red"
cursor_label = "preview_next_catalog_id" if dry_run else "next_catalog_id"
cursor_value = summary.preview_next_catalog_id if dry_run else summary.next_catalog_id
click.echo(
click.style(
f"{summary.operation} {status}. "
@@ -821,10 +848,12 @@ def _echo_bundle_archive_operation_summary(summary) -> None:
f"archive_bytes={summary.archive_bytes} duration={summary.elapsed_time:.2f}s "
f"validation_time={summary.validation_time:.2f}s "
f"runs_per_second={summary.runs_per_second:.2f} rows_per_second={summary.rows_per_second:.2f} "
f"bytes_per_second={summary.bytes_per_second:.2f}",
f"bytes_per_second={summary.bytes_per_second:.2f} {cursor_label}={cursor_value or 'none'}",
fg=fg,
)
)
if dry_run:
click.echo(click.style("Dry-run cursor is preview-only; do not persist it for a destructive run.", fg="yellow"))
click.echo(click.style("table,row_count", fg="white"))
for table_name in [
"workflow_runs",
@@ -842,7 +871,8 @@ def _echo_bundle_archive_operation_summary(summary) -> None:
click.style(
f" bundle={result.bundle_id} tenant={result.tenant_id} runs={result.run_count} "
f"rows={result.row_count} archive_bytes={result.archive_bytes} "
f"time={result.elapsed_time:.2f}s validation={result.validation_time:.2f}s",
f"catalog_id={result.catalog_id} time={result.elapsed_time:.2f}s "
f"validation={result.validation_time:.2f}s",
fg="white",
)
)
@@ -850,7 +880,7 @@ def _echo_bundle_archive_operation_summary(summary) -> None:
click.echo(
click.style(
f" failed bundle={result.bundle_id} tenant={result.tenant_id} "
f"object_prefix={result.object_prefix} error={result.error}",
f"catalog_id={result.catalog_id} object_prefix={result.object_prefix} error={result.error}",
fg="red",
)
)
@@ -867,25 +897,24 @@ def _echo_bundle_archive_operation_summary(summary) -> None:
)
@click.option("--run-id", required=False, help="Workflow run ID to restore.")
@click.option(
"--start-from",
type=click.DateTime(formats=["%Y-%m-%d", "%Y-%m-%dT%H:%M:%S"]),
"--target-month",
metavar="YYYY-MM",
default=None,
help="Optional lower bound (inclusive) for created_at; must be paired with --end-before.",
help="V2 catalog month to restore; required unless --run-id is used.",
)
@click.option(
"--end-before",
type=click.DateTime(formats=["%Y-%m-%d", "%Y-%m-%dT%H:%M:%S"]),
"--after-catalog-id",
default=None,
help="Optional upper bound (exclusive) for created_at; must be paired with --start-from.",
help="Exclusive V2 cursor from the same restore month and tenant scope.",
)
@click.option("--workers", default=1, show_default=True, type=int, help="V1 --run-id compatibility only.")
@click.option("--limit", type=int, default=100, show_default=True, help="Maximum number of V2 bundles to restore.")
@click.option("--limit", type=click.IntRange(min=1), default=100, show_default=True, help="Maximum V2 catalog rows.")
@click.option("--dry-run", is_flag=True, help="Preview without restoring.")
def restore_workflow_runs(
tenant_ids: str | None,
run_id: str | None,
start_from: datetime.datetime | None,
end_before: datetime.datetime | None,
target_month: str | None,
after_catalog_id: str | None,
workers: int,
limit: int,
dry_run: bool,
@@ -905,23 +934,20 @@ def restore_workflow_runs(
from services.retention.workflow_run.bundle_archive_maintenance import WorkflowRunBundleArchiveMaintenance
from services.retention.workflow_run.restore_archived_workflow_run import WorkflowRunRestore
parsed_tenant_ids = None
if tenant_ids:
parsed_tenant_ids = [tid.strip() for tid in tenant_ids.split(",") if tid.strip()]
if not parsed_tenant_ids:
raise click.BadParameter("tenant-ids must not be empty")
parsed_tenant_ids = _parse_comma_separated_ids(tenant_ids, param_name="tenant-ids")
if (start_from is None) ^ (end_before is None):
raise click.UsageError("--start-from and --end-before must be provided together.")
if run_id is None and (start_from is None or end_before is None):
raise click.UsageError("--start-from and --end-before are required for batch restore.")
if workers < 1:
raise click.BadParameter("workers must be at least 1")
if run_id is not None and (target_month is not None or after_catalog_id is not None):
raise click.UsageError("--target-month and --after-catalog-id are only valid for V2 batch restore.")
if run_id is None and target_month is None:
raise click.UsageError("--target-month is required for V2 batch restore.")
start_time = datetime.datetime.now(datetime.UTC)
target_desc = f"workflow run {run_id}" if run_id else f"workflow archive catalog month {target_month}"
click.echo(
click.style(
f"Starting restore of workflow run {run_id} at {start_time.isoformat()}.",
f"Starting restore of {target_desc} at {start_time.isoformat()}.",
fg="white",
)
)
@@ -955,17 +981,20 @@ def restore_workflow_runs(
click.echo(
click.style("--workers is ignored for V2 bundle restore; bundles are processed serially.", fg="yellow")
)
assert start_from is not None
assert end_before is not None
assert target_month is not None
target_year, target_month_number = _parse_archive_target_month(target_month)
catalog_cursor = _parse_archive_catalog_cursor(after_catalog_id)
bundle_restorer = WorkflowRunBundleArchiveMaintenance(dry_run=dry_run, strict_content_validation=True)
summary = bundle_restorer.restore_batch(
tenant_ids=parsed_tenant_ids,
start_date=start_from,
end_date=end_before,
target_year=target_year,
target_month=target_month_number,
after_catalog_id=catalog_cursor,
limit=limit,
)
_echo_bundle_archive_operation_summary(summary)
return
_echo_bundle_archive_operation_summary(summary, dry_run=dry_run)
if summary.bundles_failed:
raise click.exceptions.Exit(1)
@click.command(
@@ -979,23 +1008,41 @@ def restore_workflow_runs(
)
@click.option("--run-id", required=False, help="Workflow run ID to delete.")
@click.option(
"--start-from",
type=click.DateTime(formats=["%Y-%m-%d", "%Y-%m-%dT%H:%M:%S"]),
"--target-month",
metavar="YYYY-MM",
default=None,
help="Optional lower bound (inclusive) for created_at; must be paired with --end-before.",
help="V2 catalog month to delete; required unless --run-id is used.",
)
@click.option(
"--end-before",
type=click.DateTime(formats=["%Y-%m-%d", "%Y-%m-%dT%H:%M:%S"]),
"--after-catalog-id",
default=None,
help="Optional upper bound (exclusive) for created_at; must be paired with --start-from.",
help="Exclusive V2 cursor from the same delete month and tenant scope.",
)
@click.option(
"--run-shard-index",
default=None,
type=click.IntRange(min=0),
help="Zero-based archive shard index. Must be paired with --run-shard-total.",
)
@click.option(
"--run-shard-total",
default=None,
type=click.IntRange(min=1, max=16),
help="Total archive shard count. Must be paired with --run-shard-index.",
)
@click.option("--all-pages", is_flag=True, help="Process catalog pages until an empty page is reached.")
@click.option(
"--limit",
type=click.IntRange(min=1),
default=100,
show_default=True,
help="Maximum V2 catalog rows per page.",
)
@click.option("--limit", type=int, default=100, show_default=True, help="Maximum number of V2 bundles to delete.")
@click.option("--dry-run", is_flag=True, help="Preview without deleting.")
@click.option(
"--skip-bad-archives",
is_flag=True,
help="Continue batch deletion when one archive object fails validation.",
help="V1 --run-id only: continue when one archive object fails validation.",
)
@click.option(
"--restore-sample-interval",
@@ -1007,8 +1054,11 @@ def restore_workflow_runs(
def delete_archived_workflow_runs(
tenant_ids: str | None,
run_id: str | None,
start_from: datetime.datetime | None,
end_before: datetime.datetime | None,
target_month: str | None,
after_catalog_id: str | None,
run_shard_index: int | None,
run_shard_total: int | None,
all_pages: bool,
limit: int,
dry_run: bool,
skip_bad_archives: bool,
@@ -1018,26 +1068,38 @@ def delete_archived_workflow_runs(
Delete archived workflow runs from the database.
Batch delete uses V2 bundle metadata and validates object existence, manifest schema, object size, checksum, row
counts, and source/archive content checksums before deleting source rows. `--run-id` keeps the V1 per-run path.
counts, and source/archive content checksums before deleting source rows. Parallel workers may select one exact
archive shard; all-pages mode keeps only the current bounded page in memory. `--run-id` keeps the V1 per-run path.
"""
from services.retention.workflow_run.bundle_archive_maintenance import WorkflowRunBundleArchiveMaintenance
from services.retention.workflow_run.delete_archived_workflow_run import ArchivedWorkflowRunDeletion
parsed_tenant_ids = None
if tenant_ids:
parsed_tenant_ids = [tid.strip() for tid in tenant_ids.split(",") if tid.strip()]
if not parsed_tenant_ids:
raise click.BadParameter("tenant-ids must not be empty")
parsed_tenant_ids = _parse_comma_separated_ids(tenant_ids, param_name="tenant-ids")
if (start_from is None) ^ (end_before is None):
raise click.UsageError("--start-from and --end-before must be provided together.")
if run_id is None and (start_from is None or end_before is None):
raise click.UsageError("--start-from and --end-before are required for batch delete.")
if restore_sample_interval < 0:
raise click.BadParameter("restore-sample-interval must be >= 0")
if run_id is not None and (
target_month is not None
or after_catalog_id is not None
or run_shard_index is not None
or run_shard_total is not None
or all_pages
):
raise click.UsageError(
"--target-month, --after-catalog-id, --run-shard-index, --run-shard-total, and --all-pages "
"are only valid for V2 batch delete."
)
if run_id is None and target_month is None:
raise click.UsageError("--target-month is required for V2 batch delete.")
if run_id is None and skip_bad_archives:
raise click.UsageError("--skip-bad-archives is not supported for V2 catalog batches; they fail fast.")
if (run_shard_index is None) ^ (run_shard_total is None):
raise click.UsageError("--run-shard-index and --run-shard-total must be provided together.")
if run_shard_index is not None and run_shard_total is not None and run_shard_index >= run_shard_total:
raise click.UsageError("--run-shard-index must be less than --run-shard-total.")
start_time = datetime.datetime.now(datetime.UTC)
target_desc = f"workflow run {run_id}" if run_id else "workflow runs"
target_desc = f"workflow run {run_id}" if run_id else f"workflow archive catalog month {target_month}"
click.echo(
click.style(
f"Starting delete of {target_desc} at {start_time.isoformat()}.",
@@ -1110,20 +1172,104 @@ def delete_archived_workflow_runs(
if restore_sample_interval:
click.echo(click.style("--restore-sample-interval is ignored for V2 bundle delete.", fg="yellow"))
assert start_from is not None
assert end_before is not None
bundle_deleter = WorkflowRunBundleArchiveMaintenance(
dry_run=dry_run,
strict_content_validation=True,
stop_on_error=not skip_bad_archives,
assert target_month is not None
target_year, target_month_number = _parse_archive_target_month(target_month)
catalog_cursor = _parse_archive_catalog_cursor(after_catalog_id)
shard = (
f"{run_shard_index:02d}-of-{run_shard_total:02d}"
if run_shard_index is not None and run_shard_total is not None
else None
)
summary = bundle_deleter.delete_batch(
tenant_ids=parsed_tenant_ids,
start_date=start_from,
end_date=end_before,
limit=limit,
)
_echo_bundle_archive_operation_summary(summary)
bundle_deleter = WorkflowRunBundleArchiveMaintenance(dry_run=dry_run, strict_content_validation=True)
if run_shard_total is not None:
try:
bundle_deleter.validate_catalog_shards(
target_year=target_year,
target_month=target_month_number,
shard_total=run_shard_total,
tenant_ids=parsed_tenant_ids,
)
except ValueError as exc:
logger.exception(
"Archive catalog shard preflight failed: target_month=%s shard=%s",
target_month,
shard,
)
raise click.ClickException(
f"Archive catalog shard preflight failed for target_month={target_month} shard={shard}: {exc}"
) from exc
initial_catalog_cursor = catalog_cursor
pages_processed = 0
bundles_succeeded = 0
runs_processed = 0
rows_processed = 0
archive_bytes = 0
while True:
summary = bundle_deleter.delete_batch(
tenant_ids=parsed_tenant_ids,
target_year=target_year,
target_month=target_month_number,
after_catalog_id=catalog_cursor,
limit=limit,
shard=shard,
)
_echo_bundle_archive_operation_summary(summary, dry_run=dry_run)
if summary.bundles_failed:
failed_result = next((result for result in summary.results if not result.success), None)
failed_catalog_id = failed_result.catalog_id if failed_result is not None else "unknown"
page_resume_cursor = summary.preview_next_catalog_id if dry_run else summary.next_catalog_id
resume_cursor = page_resume_cursor or catalog_cursor
if dry_run:
cursor_details = (
f"preview_after_catalog_id={resume_cursor or 'none'} "
f"destructive_retry_after_catalog_id={initial_catalog_cursor or 'none'}"
)
else:
cursor_details = f"resume_after_catalog_id={resume_cursor or 'none'}"
click.echo(
click.style(
f"Delete stopped: target_month={target_month} shard={shard or 'all'} "
f"failed_catalog_id={failed_catalog_id} "
f"{cursor_details}",
fg="red",
)
)
raise click.exceptions.Exit(1)
if not all_pages:
break
if summary.bundles_processed == 0:
break
pages_processed += 1
bundles_succeeded += summary.bundles_succeeded
runs_processed += summary.runs_processed
rows_processed += summary.rows_processed
archive_bytes += summary.archive_bytes
next_catalog_id = summary.preview_next_catalog_id if dry_run else summary.next_catalog_id
if next_catalog_id is None or (catalog_cursor is not None and next_catalog_id <= catalog_cursor):
click.echo(
click.style(
f"Delete cursor did not advance: target_month={target_month} shard={shard or 'all'} "
f"after_catalog_id={catalog_cursor or 'none'} next_catalog_id={next_catalog_id or 'none'}",
fg="red",
)
)
raise click.exceptions.Exit(1)
catalog_cursor = next_catalog_id
if all_pages:
final_cursor_label = "preview_final_catalog_id" if dry_run else "final_catalog_id"
click.echo(
click.style(
f"Delete all-pages completed successfully. target_month={target_month} shard={shard or 'all'} "
f"pages={pages_processed} bundles_success={bundles_succeeded} runs={runs_processed} "
f"rows={rows_processed} archive_bytes={archive_bytes} "
f"{final_cursor_label}={catalog_cursor or 'none'}",
fg="green",
)
)
def _find_orphaned_draft_variables(batch_size: int = 1000) -> list[str]:
+6
View File
@@ -14,6 +14,12 @@ class EnterpriseFeatureConfig(BaseSettings):
default=False,
)
WEBAPP_PUBLIC_ACCESS_ENABLED: bool = Field(
description="Whether admins are allowed to set a webapp's access mode to public (anyone with the link, "
"no auth). Disable in security-sensitive on-prem deployments.",
default=True,
)
CAN_REPLACE_LOGO: bool = Field(
description="Allow customization of the enterprise logo.",
default=False,
+2
View File
@@ -1,5 +1,6 @@
from configs.extra.agent_backend_config import AgentBackendConfig
from configs.extra.archive_config import ArchiveStorageConfig
from configs.extra.knowledge_fs_config import KnowledgeFSConfig
from configs.extra.notion_config import NotionConfig
from configs.extra.sentry_config import SentryConfig
@@ -8,6 +9,7 @@ class ExtraServiceConfig(
# place the configs in alphabet order
AgentBackendConfig,
ArchiveStorageConfig,
KnowledgeFSConfig,
NotionConfig,
SentryConfig,
):
+16 -1
View File
@@ -1,4 +1,4 @@
from pydantic import Field, NonNegativeFloat
from pydantic import Field, NonNegativeFloat, NonNegativeInt, PositiveFloat
from pydantic_settings import BaseSettings
@@ -22,6 +22,21 @@ class AgentBackendConfig(BaseSettings):
default="success",
)
AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS: PositiveFloat = Field(
description="Read timeout for one Agent backend SSE connection.",
default=30,
)
AGENT_BACKEND_STREAM_MAX_RECONNECTS: NonNegativeInt = Field(
description="Maximum Agent backend SSE reconnects before failing the run.",
default=3,
)
AGENT_BACKEND_RUN_TIMEOUT_SECONDS: PositiveFloat = Field(
description="Total deadline for one Agent backend run event stream.",
default=1200,
)
AGENT_SHELL_ENABLED: bool = Field(
description=(
"Inject the dify.shell layer (sandboxed bash workspace) into Agent runs. "
+64
View File
@@ -0,0 +1,64 @@
"""Configuration for the optional KnowledgeFS Console bridge."""
from urllib.parse import urlsplit
from pydantic import Field, PositiveFloat, SecretStr, field_validator, model_validator
from pydantic_settings import BaseSettings
class KnowledgeFSConfig(BaseSettings):
"""Server-only settings for the KnowledgeFS production connection."""
KNOWLEDGE_FS_ENABLED: bool = Field(
default=False,
description="Enable the private KnowledgeFS Console bridge.",
)
KNOWLEDGE_FS_BASE_URL: str | None = Field(default=None, description="KnowledgeFS gateway base URL.")
KNOWLEDGE_FS_JWT_SECRET: SecretStr | None = Field(
default=None,
min_length=32,
description="Shared secret used to sign short-lived KnowledgeFS service JWTs.",
)
KNOWLEDGE_FS_SSE_READ_TIMEOUT_SECONDS: PositiveFloat = Field(default=300.0, le=3600.0, allow_inf_nan=False)
KNOWLEDGE_FS_TIMEOUT_SECONDS: PositiveFloat = Field(default=10.0, le=60.0, allow_inf_nan=False)
@field_validator(
"KNOWLEDGE_FS_BASE_URL",
"KNOWLEDGE_FS_JWT_SECRET",
mode="before",
)
@classmethod
def normalize_optional_string(cls, value: object) -> object:
if isinstance(value, SecretStr):
normalized = value.get_secret_value().strip()
return SecretStr(normalized) if normalized else None
if isinstance(value, str):
normalized = value.strip()
return normalized or None
return value
@field_validator("KNOWLEDGE_FS_BASE_URL")
@classmethod
def validate_base_url(cls, value: str | None) -> str | None:
if value is None:
return None
parsed = urlsplit(value)
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
raise ValueError("KNOWLEDGE_FS_BASE_URL must be an absolute HTTP(S) URL")
try:
_ = parsed.port
except ValueError as exc:
raise ValueError("KNOWLEDGE_FS_BASE_URL must include a valid port") from exc
if parsed.username or parsed.password or parsed.query or parsed.fragment:
raise ValueError("KNOWLEDGE_FS_BASE_URL must not include credentials, query, or fragment")
return value.rstrip("/")
@model_validator(mode="after")
def validate_enabled_connection(self) -> "KnowledgeFSConfig":
if not self.KNOWLEDGE_FS_ENABLED:
return self
if bool(self.KNOWLEDGE_FS_BASE_URL) != bool(self.KNOWLEDGE_FS_JWT_SECRET):
raise ValueError("KNOWLEDGE_FS_BASE_URL and KNOWLEDGE_FS_JWT_SECRET must be configured together")
if not self.KNOWLEDGE_FS_BASE_URL:
raise ValueError("KnowledgeFS connection settings are required when the integration is enabled")
return self
+74 -1
View File
@@ -1,4 +1,4 @@
from datetime import timedelta
from datetime import datetime, timedelta
from enum import StrEnum
from typing import Literal
@@ -11,6 +11,7 @@ from pydantic import (
PositiveFloat,
PositiveInt,
computed_field,
field_validator,
)
from pydantic_settings import BaseSettings
@@ -275,6 +276,63 @@ class PluginConfig(BaseSettings):
default=50 * 1024 * 1024,
)
NEW_USER_DEFAULT_PLUGIN_IDS: str = Field(
description="Comma-separated marketplace plugin IDs whose latest versions are installed for new users",
default="",
)
@field_validator("PLUGIN_REMOTE_INSTALL_PORT", mode="before")
@classmethod
def _reject_host_port_shaped_plugin_remote_install_port(cls, v):
"""Reject ``host:port``-shaped values with an actionable hint.
``EXPOSE_PLUGIN_DEBUGGING_PORT`` is overloaded: it feeds both the
plugin_daemon ``ports:`` mapping (where ``127.0.0.1:5003`` is valid
compose syntax) and this integer app setting advertised in the console.
Without this guard a loopback bind spec crashloops the api container
with an opaque ``int_parsing`` traceback. See issue #39323.
"""
if isinstance(v, str) and ":" in v.strip():
raise ValueError(
"PLUGIN_REMOTE_INSTALL_PORT must be a bare port number, got "
f"{v!r}. A 'host:port' value usually means "
"EXPOSE_PLUGIN_DEBUGGING_PORT was set to a compose publish spec "
"like '127.0.0.1:5003'; bind loopback via a "
"docker-compose.override.yaml instead of overloading this var."
)
return v
@property
def NEW_USER_DEFAULT_PLUGIN_ID_LIST(self) -> list[str]:
return [item.strip() for item in self.NEW_USER_DEFAULT_PLUGIN_IDS.split(",") if item.strip()]
NEW_USER_DEFAULT_MODELS: str = Field(
description=("Comma-separated default models for new users in 'model_type:provider:model' format"),
default="",
)
@property
def NEW_USER_DEFAULT_MODEL_LIST(self) -> list[tuple[str, str, str]]:
default_models: list[tuple[str, str, str]] = []
configured_model_types: set[str] = set()
for item in self.NEW_USER_DEFAULT_MODELS.split(","):
if not item.strip():
continue
parts = tuple(part.strip() for part in item.split(":", 2))
if len(parts) != 3 or not all(parts):
raise ValueError("NEW_USER_DEFAULT_MODELS entries must use 'model_type:provider:model' format")
model_type, provider, model = parts
if model_type in configured_model_types:
raise ValueError(f"NEW_USER_DEFAULT_MODELS contains duplicate model type: {model_type}")
configured_model_types.add(model_type)
default_models.append((model_type, provider, model))
return default_models
class MarketplaceConfig(BaseSettings):
"""
@@ -784,6 +842,11 @@ class WorkflowConfig(BaseSettings):
default=500,
)
WORKFLOW_GENERATOR_NODE_BUILDER_MAX_WORKERS: PositiveInt = Field(
description="Maximum concurrent node-builder LLM calls per workflow generation request",
default=6,
)
WORKFLOW_MAX_EXECUTION_TIME: PositiveInt = Field(
description="Maximum execution time in seconds for a single workflow",
default=1200,
@@ -1097,6 +1160,16 @@ class HomepageConfig(BaseSettings):
default=True,
)
ENABLE_STEP_BY_STEP_TOUR: bool = Field(
description="Enable account-level Step-by-step Tour eligibility checks",
default=False,
)
STEP_BY_STEP_TOUR_ROLLOUT_STARTED_AT: datetime | None = Field(
description="UTC timestamp after which newly initialized accounts are eligible for Step-by-step Tour",
default=None,
)
class RagEtlConfig(BaseSettings):
"""
+4
View File
@@ -38,7 +38,9 @@ from . import (
feature,
human_input_form,
init_validate,
knowledge_fs_proxy,
notification,
onboarding,
ping,
setup,
spec,
@@ -195,6 +197,7 @@ __all__ = [
"human_input_form",
"init_validate",
"installed_app",
"knowledge_fs_proxy",
"load_balancing_config",
"login",
"mcp_server",
@@ -207,6 +210,7 @@ __all__ = [
"notification",
"oauth",
"oauth_server",
"onboarding",
"ops_trace",
"parameter",
"ping",
+2
View File
@@ -105,6 +105,7 @@ class SyncDraftWorkflowPayload(BaseModel):
graph: dict[str, Any]
features: dict[str, Any]
hash: str | None = None
is_collaborative: bool = Field(default=False, alias="_is_collaborative")
environment_variables: list[dict[str, Any]] = Field(
default_factory=list,
)
@@ -610,6 +611,7 @@ class DraftWorkflowApi(Resource):
environment_variables=environment_variables,
conversation_variables=conversation_variables,
session=db.session(),
graph_only=args["is_collaborative"],
)
except WorkflowHashNotEqualError:
raise DraftWorkflowNotSync()
+8 -1
View File
@@ -607,6 +607,13 @@ class DatasetListApi(Resource):
ReplaceMemberBindings(scope=RBACResourceWhitelistScope.ALL),
)
initialize_created_app_rbac_access_task.delay(current_tenant_id, current_user.id, dataset_id=dataset.id)
else:
enterprise_rbac_service.RBACService.DatasetAccess.replace_whitelist(
current_tenant_id,
current_user.id,
dataset.id,
ReplaceMemberBindings(scope=RBACResourceWhitelistScope.SPECIFIC),
)
permission_keys_map = enterprise_rbac_service.RBACService.DatasetPermissions.batch_get(
current_tenant_id,
@@ -875,7 +882,7 @@ class DatasetIndexingEstimateApi(Resource):
file_details = session.scalars(
select(UploadFile).where(UploadFile.tenant_id == current_tenant_id, UploadFile.id.in_(file_ids))
).all()
if file_details is None:
if not file_details:
raise NotFound("File not found.")
if file_details:
+48 -3
View File
@@ -47,7 +47,9 @@ from controllers.console.explore.error import (
NotWorkflowAppError,
)
from controllers.console.explore.wraps import TrialAppResource, trial_feature_enable
from controllers.console.wraps import with_current_user
from controllers.console.files import FILE_UPLOAD_PARAMS, upload_file_from_request
from controllers.console.remote_files import RemoteFileUploadPayload, upload_remote_file_from_request
from controllers.console.wraps import cloud_edition_billing_resource_check, with_current_user
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict
from core.app.apps.base_app_queue_manager import AppQueueManager
@@ -61,12 +63,13 @@ from extensions.ext_database import db
from extensions.ext_redis import redis_client
from fields.base import ResponseModel
from fields.conversation_variable_fields import WorkflowConversationVariableResponse
from fields.file_fields import FileResponse, FileWithSignedUrl
from fields.message_fields import SuggestedQuestionsResponse
from graphon.graph_engine.manager import GraphEngineManager
from graphon.model_runtime.errors.invoke import InvokeError
from libs import helper
from libs.helper import dump_response, to_timestamp, uuid_value
from models import Account
from models import Account, App
from models.account import TenantStatus
from models.model import AppMode, Site, load_annotation_reply_config
from models.workflow import Workflow
@@ -428,6 +431,36 @@ register_response_schema_models(
simple_account_model = console_ns.models[TrialSimpleAccount.__name__]
class TrialAppFileUploadApi(TrialAppResource):
@trial_feature_enable
@cloud_edition_billing_resource_check("documents")
@console_ns.doc(consumes=["multipart/form-data"], params=FILE_UPLOAD_PARAMS)
@console_ns.response(201, "File uploaded successfully", console_ns.models[FileResponse.__name__])
@with_current_user
def post(self, current_user: Account, app_model: App):
"""Upload a file into the tenant that owns the trial app."""
upload_file = upload_file_from_request(
current_user=current_user,
resource_tenant_id=app_model.tenant_id,
)
return dump_response(FileResponse, upload_file), 201
class TrialAppRemoteFileUploadApi(TrialAppResource):
@trial_feature_enable
@cloud_edition_billing_resource_check("documents")
@console_ns.expect(console_ns.models[RemoteFileUploadPayload.__name__])
@console_ns.response(201, "File uploaded successfully", console_ns.models[FileWithSignedUrl.__name__])
@with_current_user
def post(self, current_user: Account, app_model: App):
"""Upload a remote file into the tenant that owns the trial app."""
remote_file = upload_remote_file_from_request(
current_user=current_user,
resource_tenant_id=app_model.tenant_id,
)
return remote_file.model_dump(mode="json"), 201
class TrialAppWorkflowRunApi(TrialAppResource):
@trial_feature_enable
@console_ns.expect(console_ns.models[WorkflowRunRequest.__name__])
@@ -612,7 +645,7 @@ class TrialChatAudioApi(TrialAppResource):
def post(self, current_user: Account, trial_app):
app_model = trial_app
file = request.files["file"]
file = request.files.get("file")
try:
# Get IDs before they might be detached from session
@@ -887,6 +920,18 @@ class DatasetListApi(Resource):
console_ns.add_resource(TrialChatApi, "/trial-apps/<uuid:app_id>/chat-messages", endpoint="trial_app_chat_completion")
console_ns.add_resource(
TrialAppFileUploadApi,
"/trial-apps/<uuid:app_id>/files/upload",
endpoint="trial_app_file_upload",
)
console_ns.add_resource(
TrialAppRemoteFileUploadApi,
"/trial-apps/<uuid:app_id>/remote-files/upload",
endpoint="trial_app_remote_file_upload",
)
console_ns.add_resource(
TrialMessageSuggestedQuestionApi,
"/trial-apps/<uuid:app_id>/messages/<uuid:message_id>/suggested-questions",
+41 -35
View File
@@ -29,7 +29,7 @@ from extensions.ext_database import db
from fields.file_fields import FileResponse, UploadConfig
from libs.helper import dump_response
from libs.login import login_required
from models import Account
from models import Account, UploadFile
from services.file_service import FileService
from . import console_ns
@@ -39,7 +39,7 @@ register_response_schema_models(console_ns, AllowedExtensionsResponse, TextConte
PREVIEW_WORDS_LIMIT = 3000
_FILE_UPLOAD_PARAMS = {
FILE_UPLOAD_PARAMS = {
"file": {
"description": "File to upload",
"in": "formData",
@@ -56,6 +56,43 @@ _FILE_UPLOAD_PARAMS = {
}
def upload_file_from_request(*, current_user: Account, resource_tenant_id: str | None = None) -> UploadFile:
"""Validate the multipart request and persist the file under the requested resource tenant."""
source_str = request.form.get("source")
source: Literal["datasets"] | None = "datasets" if source_str == "datasets" else None
if "file" not in request.files:
raise NoFileUploadedError()
if len(request.files) > 1:
raise TooManyFilesError()
file = request.files["file"]
if not file.filename:
raise FilenameNotExistsError
if source == "datasets" and not current_user.is_dataset_editor:
raise Forbidden()
if source not in ("datasets", None):
source = None
try:
return FileService(db.engine).upload_file(
filename=file.filename,
content=file.stream.read(),
mimetype=file.mimetype,
user=current_user,
tenant_id=resource_tenant_id,
source=source,
)
except services.errors.file.FileTooLargeError as file_too_large_error:
raise FileTooLargeError(file_too_large_error.description)
except services.errors.file.UnsupportedFileTypeError:
raise UnsupportedFileTypeError()
except services.errors.file.BlockedFileExtensionError as blocked_extension_error:
raise BlockedFileExtensionError(blocked_extension_error.description)
@console_ns.route("/files/upload")
class FileApi(Resource):
@setup_required
@@ -81,42 +118,11 @@ class FileApi(Resource):
@login_required
@account_initialization_required
@cloud_edition_billing_resource_check("documents")
@console_ns.doc(consumes=["multipart/form-data"], params=_FILE_UPLOAD_PARAMS)
@console_ns.doc(consumes=["multipart/form-data"], params=FILE_UPLOAD_PARAMS)
@console_ns.response(201, "File uploaded successfully", console_ns.models[FileResponse.__name__])
@with_current_user
def post(self, current_user: Account):
source_str = request.form.get("source")
source: Literal["datasets"] | None = "datasets" if source_str == "datasets" else None
if "file" not in request.files:
raise NoFileUploadedError()
if len(request.files) > 1:
raise TooManyFilesError()
file = request.files["file"]
if not file.filename:
raise FilenameNotExistsError
if source == "datasets" and not current_user.is_dataset_editor:
raise Forbidden()
if source not in ("datasets", None):
source = None
try:
upload_file = FileService(db.engine).upload_file(
filename=file.filename,
content=file.stream.read(),
mimetype=file.mimetype,
user=current_user,
source=source,
)
except services.errors.file.FileTooLargeError as file_too_large_error:
raise FileTooLargeError(file_too_large_error.description)
except services.errors.file.UnsupportedFileTypeError:
raise UnsupportedFileTypeError()
except services.errors.file.BlockedFileExtensionError as blocked_extension_error:
raise BlockedFileExtensionError(blocked_extension_error.description)
upload_file = upload_file_from_request(current_user=current_user)
return dump_response(FileResponse, upload_file), 201
@@ -0,0 +1,332 @@
"""Authenticated transport adapter for the Console-to-KnowledgeFS proxy.
These raw Blueprint routes deliberately stay outside Dify's OpenAPI surface:
KnowledgeFS owns the wire contract consumed by the frontend. The catch-all path
avoids resource-specific Dify controllers, while the forwarding module consumes
only the operations explicitly enabled by Dify's product registry. The registry
can be validated explicitly against the pinned KnowledgeFS contract during development.
Console auth and contract-specific dataset RBAC run before forwarding. Request
bodies are capped at 64 MiB, JSON and binary responses have separate bounds,
SSE responses remain streaming with a bounded idle read timeout, and only safe
response headers are exposed. Upstream 401 responses become 502 so they cannot
trigger Dify browser-session recovery; resource-level 403 responses remain 403.
"""
from __future__ import annotations
import logging
from collections.abc import Callable, Iterator
from functools import wraps
from http import HTTPStatus
from typing import NoReturn, cast
import httpx
from flask import Response, request, stream_with_context
from flask.typing import ResponseReturnValue
from werkzeug.exceptions import (
BadGateway,
Forbidden,
GatewayTimeout,
NotFound,
RequestEntityTooLarge,
ServiceUnavailable,
)
from configs import dify_config
from controllers.console import api, bp
from controllers.console.wraps import (
account_initialization_required,
cloud_edition_billing_rate_limit_check,
setup_required,
)
from core.helper import ssrf_proxy
from libs.login import current_account_with_tenant, login_required
from services.knowledge_fs_proxy import (
KnowledgeFSAccessDeniedError,
KnowledgeFSConfigurationError,
KnowledgeFSMethod,
KnowledgeFSRouteNotAllowedError,
KnowledgeFSTimeoutError,
KnowledgeFSTransportError,
KnowledgeFSUpstreamResponse,
authorize_knowledge_fs_request,
get_knowledge_fs_operation,
proxy_knowledge_fs_request,
)
logger = logging.getLogger(__name__)
_MAX_PROXY_BODY_BYTES = 64 * 1024 * 1024
_RESPONSE_HEADER_ALLOWLIST = (
"Cache-Control",
"Content-Disposition",
"Content-Type",
"Retry-After",
"X-Trace-Id",
)
_RESPONSE_HEADER_DENYLIST = frozenset(
{
"authorization",
"connection",
"cookie",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"set-cookie",
"te",
"trailer",
"transfer-encoding",
"upgrade",
}
)
def _console_api_errors[**P](
view: Callable[P, ResponseReturnValue],
) -> Callable[P, ResponseReturnValue]:
"""Route raw Blueprint exceptions through the Console API JSON handlers."""
@wraps(view)
def decorated(*args: P.args, **kwargs: P.kwargs) -> ResponseReturnValue:
try:
return view(*args, **kwargs)
except Exception as exc:
return api.handle_error(exc)
return decorated
def _knowledge_fs_enabled[**P](
view: Callable[P, ResponseReturnValue],
) -> Callable[P, ResponseReturnValue]:
"""Hide the complete KnowledgeFS route surface while the bridge is disabled."""
@wraps(view)
def decorated(*args: P.args, **kwargs: P.kwargs) -> ResponseReturnValue:
if not dify_config.KNOWLEDGE_FS_ENABLED:
raise NotFound()
return view(*args, **kwargs)
return decorated
def _translate_proxy_error(exc: Exception, *, tenant_id: str) -> NoReturn:
"""Map forwarding failures to the stable Console HTTP error surface."""
if isinstance(exc, KnowledgeFSRouteNotAllowedError):
raise NotFound() from exc
if isinstance(exc, KnowledgeFSAccessDeniedError):
raise Forbidden() from exc
if isinstance(exc, KnowledgeFSConfigurationError):
logger.error("KnowledgeFS request was blocked by invalid configuration for tenant_id=%s", tenant_id)
raise ServiceUnavailable("KnowledgeFS integration is misconfigured") from exc
if isinstance(exc, KnowledgeFSTimeoutError):
raise GatewayTimeout("KnowledgeFS request timed out") from exc
if isinstance(exc, KnowledgeFSTransportError):
logger.warning("KnowledgeFS transport request failed for tenant_id=%s", tenant_id)
raise BadGateway("KnowledgeFS is unavailable") from exc
raise exc
def _knowledge_fs_operation_access_required(
view: Callable[[KnowledgeFSMethod, str], ResponseReturnValue],
) -> Callable[[KnowledgeFSMethod, str], ResponseReturnValue]:
"""Authorize one declared operation before billing and request-body work."""
@wraps(view)
def decorated(method: KnowledgeFSMethod, upstream_path: str) -> ResponseReturnValue:
try:
operation = get_knowledge_fs_operation(method, upstream_path)
except KnowledgeFSRouteNotAllowedError as exc:
raise NotFound() from exc
current_user, tenant_id = current_account_with_tenant()
try:
authorize_knowledge_fs_request(
account=current_user,
tenant_id=tenant_id,
operation=operation,
)
except KnowledgeFSAccessDeniedError as exc:
_translate_proxy_error(exc, tenant_id=tenant_id)
return view(method, upstream_path)
return decorated
def _request_body() -> bytes:
"""Read the raw body up to the proxy limit or raise RequestEntityTooLarge."""
body = request.stream.read(_MAX_PROXY_BODY_BYTES + 1)
if len(body) > _MAX_PROXY_BODY_BYTES:
raise RequestEntityTooLarge("KnowledgeFS proxy request body is too large")
return body
def _stream_response_body(
upstream: httpx.Response,
*,
tenant_id: str,
max_response_bytes: int,
) -> Iterator[bytes]:
"""Yield one bounded SSE response and always release its pooled connection."""
total_bytes = 0
try:
for chunk in upstream.iter_bytes():
total_bytes += len(chunk)
if total_bytes > max_response_bytes:
logger.warning("KnowledgeFS stream exceeded the proxy limit for tenant_id=%s", tenant_id)
raise ssrf_proxy.ResponseTooLargeError(f"response exceeded {max_response_bytes} bytes")
yield chunk
finally:
upstream.close()
def _proxy_response(
upstream_result: KnowledgeFSUpstreamResponse,
*,
tenant_id: str,
contract_response_headers: tuple[str, ...],
max_response_bytes: int,
) -> Response:
"""Expose raw content, status, and allowlisted headers from KnowledgeFS.
Raises:
BadGateway: KnowledgeFS rejects the configured server credential.
Forbidden: KnowledgeFS denies the account access to the requested resource.
"""
upstream = upstream_result.response
if upstream.status_code == HTTPStatus.UNAUTHORIZED:
upstream.close()
logger.error(
"KnowledgeFS rejected the Dify server credential with HTTP %s for tenant_id=%s",
upstream.status_code,
tenant_id,
)
raise BadGateway("KnowledgeFS authentication failed")
if upstream.status_code == HTTPStatus.FORBIDDEN:
upstream.close()
raise Forbidden()
allowed_header_names = dict.fromkeys(
name.lower() for name in (*_RESPONSE_HEADER_ALLOWLIST, *contract_response_headers)
)
headers = {
name: value
for name in allowed_header_names
if name not in _RESPONSE_HEADER_DENYLIST
if (value := upstream.headers.get(name)) is not None
}
if upstream_result.response_kind == "stream":
response = Response(
stream_with_context( # pyrefly: ignore[no-matching-overload]
_stream_response_body(
upstream,
tenant_id=tenant_id,
max_response_bytes=max_response_bytes,
)
),
status=upstream.status_code,
headers=headers,
)
response.call_on_close(upstream.close)
return response
try:
content = upstream.content
finally:
upstream.close()
return Response(content, status=upstream.status_code, headers=headers)
def _proxy_request(method: KnowledgeFSMethod, upstream_path: str) -> Response:
"""Forward the current raw request and return its filtered upstream response.
The call performs one outbound KnowledgeFS request. Integration failures are
converted to Console HTTP exceptions for the outer JSON error adapter.
"""
if not dify_config.KNOWLEDGE_FS_ENABLED:
raise NotFound()
current_user, tenant_id = current_account_with_tenant()
try:
proxy_result = proxy_knowledge_fs_request(
account=current_user,
method=method,
path=upstream_path,
tenant_id=tenant_id,
accept=request.headers.get("Accept"),
content_type=request.content_type,
query=request.query_string or None,
body=_request_body() if method != "GET" else None,
request_headers=request.headers,
)
except (
KnowledgeFSConfigurationError,
KnowledgeFSAccessDeniedError,
KnowledgeFSRouteNotAllowedError,
KnowledgeFSTimeoutError,
KnowledgeFSTransportError,
) as exc:
_translate_proxy_error(exc, tenant_id=tenant_id)
return _proxy_response(
proxy_result,
tenant_id=tenant_id,
contract_response_headers=proxy_result.operation.response_headers,
max_response_bytes=proxy_result.operation.max_response_bytes,
)
@_knowledge_fs_enabled
@_knowledge_fs_operation_access_required
@cloud_edition_billing_rate_limit_check("knowledge")
def _proxy_knowledge_fs_non_get(
method: KnowledgeFSMethod,
upstream_path: str,
) -> ResponseReturnValue:
"""Apply knowledge billing checks to one allowlisted non-GET operation."""
return _proxy_request(method, upstream_path)
@bp.route(
"/knowledge-fs/<path:upstream_path>",
methods=["GET", "OPTIONS"],
provide_automatic_options=False,
)
@_console_api_errors
@_knowledge_fs_enabled
@setup_required
@login_required
@account_initialization_required
def proxy_knowledge_fs_get(upstream_path: str) -> ResponseReturnValue:
"""Forward one authenticated, dataset-readable GET request.
Args:
upstream_path: Relative KFS path captured after the Console proxy prefix.
Returns:
The filtered raw KnowledgeFS response or a Console JSON error response.
"""
if request.method != "GET":
raise NotFound()
return _proxy_request("GET", upstream_path)
@bp.route(
"/knowledge-fs/<path:upstream_path>",
methods=["DELETE", "PATCH", "POST", "PUT"],
provide_automatic_options=False,
)
@_console_api_errors
@_knowledge_fs_enabled
@setup_required
@login_required
@account_initialization_required
def proxy_knowledge_fs_write(upstream_path: str) -> ResponseReturnValue:
"""Forward one authenticated non-GET request under its contract access policy.
Args:
upstream_path: Relative KFS path captured after the Console proxy prefix.
Returns:
The filtered raw KnowledgeFS response or a Console JSON error response.
"""
method = cast(KnowledgeFSMethod, request.method)
return _proxy_knowledge_fs_non_get(method, upstream_path)
+106
View File
@@ -0,0 +1,106 @@
"""Console onboarding APIs.
This module keeps Step-by-step Tour persistence account-scoped. Workspace IDs
are accepted only as presentation overrides; UI-only state such as minimized
panels or the currently active task stays on the frontend. PATCH requests are
action-based so callers do not replace server-side arrays with stale snapshots.
"""
from datetime import datetime
from typing import Literal, cast
from flask_restx import Resource
from pydantic import BaseModel, ConfigDict, Field, model_validator
from controllers.common.schema import register_response_schema_models, register_schema_models
from extensions.ext_database import db
from fields.base import ResponseModel
from libs.helper import dump_response
from libs.login import login_required
from models import Account
from services.step_by_step_tour_service import StepByStepTourPatch, StepByStepTourService
from . import console_ns
from .wraps import account_initialization_required, setup_required, with_current_tenant_id, with_current_user
StepByStepTourAction = Literal[
"skip",
"complete_task",
"uncomplete_task",
"enable_current_workspace",
"disable_current_workspace",
]
StepByStepTourTaskId = Literal["home", "studio", "knowledge", "integration"]
class StepByStepTourStatePatchPayload(BaseModel):
action: StepByStepTourAction = Field(description="State update action")
task_id: StepByStepTourTaskId | None = Field(default=None, description="Task ID for task actions")
model_config = ConfigDict(extra="forbid")
@model_validator(mode="after")
def validate_patch_shape(self) -> "StepByStepTourStatePatchPayload":
task_actions = {"complete_task", "uncomplete_task"}
if self.action in task_actions and self.task_id is None:
raise ValueError("task_id is required for task actions")
if self.action not in task_actions and self.task_id is not None:
raise ValueError("task_id is only supported for task actions")
return self
class StepByStepTourStateResponse(ResponseModel):
first_workspace_id: str | None = None
skipped: bool = False
completed_task_ids: list[StepByStepTourTaskId] = Field(default_factory=list)
manually_enabled_workspace_ids: list[str] = Field(default_factory=list)
manually_disabled_workspace_ids: list[str] = Field(default_factory=list)
updated_at: datetime | None = None
register_schema_models(console_ns, StepByStepTourStatePatchPayload)
register_response_schema_models(console_ns, StepByStepTourStateResponse)
@console_ns.route("/onboarding/step-by-step-tour/state")
class StepByStepTourStateApi(Resource):
@console_ns.doc("get_step_by_step_tour_state")
@console_ns.doc(description="Get account-level Step-by-step Tour state")
@console_ns.response(200, "Success", console_ns.models[StepByStepTourStateResponse.__name__])
@setup_required
@login_required
@account_initialization_required
@with_current_user
@with_current_tenant_id
def get(self, current_tenant_id: str, current_user: Account):
return dump_response(
StepByStepTourStateResponse,
StepByStepTourService.get_state(
account=current_user,
current_tenant_id=current_tenant_id,
session=db.session,
),
)
@console_ns.doc("patch_step_by_step_tour_state")
@console_ns.doc(description="Update account-level Step-by-step Tour state")
@console_ns.expect(console_ns.models[StepByStepTourStatePatchPayload.__name__])
@console_ns.response(200, "Success", console_ns.models[StepByStepTourStateResponse.__name__])
@setup_required
@login_required
@account_initialization_required
@with_current_user
@with_current_tenant_id
def patch(self, current_tenant_id: str, current_user: Account):
payload = StepByStepTourStatePatchPayload.model_validate(console_ns.payload or {})
patch = cast(StepByStepTourPatch, payload.model_dump(exclude_unset=True, exclude_none=True))
return dump_response(
StepByStepTourStateResponse,
StepByStepTourService.patch_state(
account=current_user,
current_tenant_id=current_tenant_id,
patch=patch,
session=db.session,
),
)
+57 -47
View File
@@ -46,6 +46,61 @@ class GetRemoteFileInfo(Resource):
).model_dump(mode="json")
def upload_remote_file_from_request(
*,
current_user: Account,
resource_tenant_id: str | None = None,
) -> FileWithSignedUrl:
"""Validate the JSON request, fetch its remote file, and persist it under the requested tenant."""
payload = RemoteFileUploadPayload.model_validate(console_ns.payload)
url = payload.url
# Try to fetch remote file metadata/content first
try:
resp = remote_fetcher.make_request("HEAD", url=url)
if resp.status_code != httpx.codes.OK:
resp = remote_fetcher.make_request("GET", url=url, timeout=3, follow_redirects=True)
if resp.status_code != httpx.codes.OK:
# Normalize into a user-friendly error message expected by tests
raise RemoteFileUploadError(f"Failed to fetch file from {url}: {resp.text}")
except httpx.RequestError as e:
raise RemoteFileUploadError(f"Failed to fetch file from {url}: {str(e)}")
file_info = helpers.guess_file_info_from_response(resp)
# Enforce file size limit with 400 (Bad Request) per tests' expectation
if not FileService.is_file_size_within_limit(extension=file_info.extension, file_size=file_info.size):
raise FileTooLargeError()
# Load content if needed
content = resp.content if resp.request.method == "GET" else remote_fetcher.make_request("GET", url).content
try:
upload_file = FileService(db.engine).upload_file(
filename=file_info.filename,
content=content,
mimetype=file_info.mimetype,
user=current_user,
tenant_id=resource_tenant_id,
source_url=url,
)
except services.errors.file.FileTooLargeError as file_too_large_error:
raise FileTooLargeError(file_too_large_error.description)
except services.errors.file.UnsupportedFileTypeError:
raise UnsupportedFileTypeError()
return FileWithSignedUrl(
id=upload_file.id,
name=upload_file.name,
size=upload_file.size,
extension=upload_file.extension,
url=file_helpers.get_signed_file_url(upload_file_id=upload_file.id),
mime_type=upload_file.mime_type,
created_by=upload_file.created_by,
created_at=int(upload_file.created_at.timestamp()),
)
@console_ns.route("/remote-files/upload")
class RemoteFileUpload(Resource):
@console_ns.expect(console_ns.models[RemoteFileUploadPayload.__name__])
@@ -53,53 +108,8 @@ class RemoteFileUpload(Resource):
@login_required
@with_current_user
def post(self, current_user: Account):
payload = RemoteFileUploadPayload.model_validate(console_ns.payload)
url = payload.url
# Try to fetch remote file metadata/content first
try:
resp = remote_fetcher.make_request("HEAD", url=url)
if resp.status_code != httpx.codes.OK:
resp = remote_fetcher.make_request("GET", url=url, timeout=3, follow_redirects=True)
if resp.status_code != httpx.codes.OK:
# Normalize into a user-friendly error message expected by tests
raise RemoteFileUploadError(f"Failed to fetch file from {url}: {resp.text}")
except httpx.RequestError as e:
raise RemoteFileUploadError(f"Failed to fetch file from {url}: {str(e)}")
file_info = helpers.guess_file_info_from_response(resp)
# Enforce file size limit with 400 (Bad Request) per tests' expectation
if not FileService.is_file_size_within_limit(extension=file_info.extension, file_size=file_info.size):
raise FileTooLargeError()
# Load content if needed
content = resp.content if resp.request.method == "GET" else remote_fetcher.make_request("GET", url).content
try:
upload_file = FileService(db.engine).upload_file(
filename=file_info.filename,
content=content,
mimetype=file_info.mimetype,
user=current_user,
source_url=url,
)
except services.errors.file.FileTooLargeError as file_too_large_error:
raise FileTooLargeError(file_too_large_error.description)
except services.errors.file.UnsupportedFileTypeError:
raise UnsupportedFileTypeError()
# Success: return created resource with 201 status
remote_file = upload_remote_file_from_request(current_user=current_user)
return (
FileWithSignedUrl(
id=upload_file.id,
name=upload_file.name,
size=upload_file.size,
extension=upload_file.extension,
url=file_helpers.get_signed_file_url(upload_file_id=upload_file.id),
mime_type=upload_file.mime_type,
created_by=upload_file.created_by,
created_at=int(upload_file.created_at.timestamp()),
).model_dump(mode="json"),
remote_file.model_dump(mode="json"),
201,
)
@@ -98,6 +98,7 @@ def handle_collaboration_event(sid, data):
6. workflow_update
7. comments_update
8. node_panel_presence
9. graph_view_state (session reports tab visibility; drives leader election)
"""
return collaboration_service.relay_collaboration_event(sid, data)
+10 -27
View File
@@ -4,20 +4,16 @@ from http import HTTPStatus
from flask import redirect
from flask_restx import Resource
from pydantic import BaseModel, Field
from werkzeug.exceptions import Conflict, NotFound
from werkzeug.exceptions import Conflict, Forbidden, NotFound
from controllers.common.fields import RedirectResponse
from controllers.common.schema import register_response_schema_models, register_schema_models
from controllers.console import console_ns
from controllers.console.wraps import (
RBACPermission,
RBACResourceScope,
account_initialization_required,
cloud_edition_billing_enabled,
cloud_edition_billing_paid_plan_required,
is_admin_or_owner_required,
only_edition_cloud,
rbac_permission_required,
setup_required,
)
from extensions.ext_database import db
@@ -25,6 +21,7 @@ from fields.base import ResponseModel
from libs.archive_storage import get_export_storage
from libs.helper import dump_response
from libs.login import current_account_with_tenant, login_required
from models import TenantAccountRole
from services.retention.workflow_run.archive_download_preparation import ARCHIVE_DOWNLOAD_MIME_TYPE
from services.retention.workflow_run.archive_download_task_cache import (
WorkflowRunArchiveDownloadStatus,
@@ -98,11 +95,13 @@ register_response_schema_models(
)
def _current_ids() -> tuple[str, str]:
"""Return current `(tenant_id, account_id)` or raise when no workspace is selected."""
def _current_owner_or_admin_ids() -> tuple[str, str]:
"""Return current Cloud workspace IDs for an owner or admin, independently of enterprise RBAC."""
current_user, current_tenant_id = current_account_with_tenant()
if not current_tenant_id:
raise NotFound("Current workspace not found")
if not TenantAccountRole.is_privileged_role(current_user.current_role):
raise Forbidden()
return current_tenant_id, current_user.id
@@ -124,12 +123,8 @@ class WorkflowRunArchivesApi(Resource):
@only_edition_cloud
@cloud_edition_billing_enabled
@cloud_edition_billing_paid_plan_required
@is_admin_or_owner_required
@rbac_permission_required(
RBACResourceScope.WORKSPACE, RBACPermission.WORKSPACE_ROLE_MANAGE, resource_required=False
)
def get(self):
tenant_id, _ = _current_ids()
tenant_id, _ = _current_owner_or_admin_ids()
return dump_response(WorkflowRunArchiveListResponse, list_workflow_run_archives(db.session(), tenant_id))
@@ -149,12 +144,8 @@ class WorkflowRunArchiveDownloadsApi(Resource):
@only_edition_cloud
@cloud_edition_billing_enabled
@cloud_edition_billing_paid_plan_required
@is_admin_or_owner_required
@rbac_permission_required(
RBACResourceScope.WORKSPACE, RBACPermission.WORKSPACE_ROLE_MANAGE, resource_required=False
)
def post(self):
tenant_id, account_id = _current_ids()
tenant_id, account_id = _current_owner_or_admin_ids()
payload = WorkflowRunArchiveDownloadPayload.model_validate(console_ns.payload or {})
try:
task = create_workflow_run_archive_download_task(
@@ -180,12 +171,8 @@ class WorkflowRunArchiveDownloadApi(Resource):
@only_edition_cloud
@cloud_edition_billing_enabled
@cloud_edition_billing_paid_plan_required
@is_admin_or_owner_required
@rbac_permission_required(
RBACResourceScope.WORKSPACE, RBACPermission.WORKSPACE_ROLE_MANAGE, resource_required=False
)
def get(self, download_id: str):
tenant_id, _ = _current_ids()
tenant_id, _ = _current_owner_or_admin_ids()
try:
task = get_workflow_run_archive_download_task(tenant_id=tenant_id, download_id=download_id)
except WorkflowRunArchiveDownloadTaskNotFoundError as exc:
@@ -209,12 +196,8 @@ class WorkflowRunArchiveDownloadFileApi(Resource):
@only_edition_cloud
@cloud_edition_billing_enabled
@cloud_edition_billing_paid_plan_required
@is_admin_or_owner_required
@rbac_permission_required(
RBACResourceScope.WORKSPACE, RBACPermission.WORKSPACE_ROLE_MANAGE, resource_required=False
)
def get(self, download_id: str):
tenant_id, _ = _current_ids()
tenant_id, _ = _current_owner_or_admin_ids()
try:
task = get_ready_workflow_run_archive_download_task(tenant_id=tenant_id, download_id=download_id)
except WorkflowRunArchiveDownloadTaskNotFoundError as exc:
+1 -1
View File
@@ -101,7 +101,7 @@ class AudioApi(Resource):
Accepts an audio file upload and returns the transcribed text.
"""
file = request.files["file"]
file = request.files.get("file")
try:
response = AudioService.transcript_asr(
+16 -8
View File
@@ -531,14 +531,22 @@ class DatasetListApi(DatasetApiResource):
except services.errors.dataset.DatasetNameDuplicateError:
raise DatasetNameDuplicateError()
if payload.permission == DatasetPermissionEnum.ALL_TEAM and dify_config.RBAC_ENABLED:
RBACService.DatasetAccess.replace_whitelist(
tenant_id,
current_user.id,
dataset.id,
ReplaceMemberBindings(scope=RBACResourceWhitelistScope.ALL),
)
initialize_created_app_rbac_access_task.delay(tenant_id, current_user.id, dataset_id=dataset.id)
if dify_config.RBAC_ENABLED:
if payload.permission == DatasetPermissionEnum.ALL_TEAM:
RBACService.DatasetAccess.replace_whitelist(
tenant_id,
current_user.id,
dataset.id,
ReplaceMemberBindings(scope=RBACResourceWhitelistScope.ALL),
)
initialize_created_app_rbac_access_task.delay(tenant_id, current_user.id, dataset_id=dataset.id)
else:
RBACService.DatasetAccess.replace_whitelist(
tenant_id,
current_user.id,
dataset.id,
ReplaceMemberBindings(scope=RBACResourceWhitelistScope.SPECIFIC),
)
return _dump_service_dataset_detail(dataset, session=session), 200
+1 -1
View File
@@ -76,7 +76,7 @@ class AudioApi(WebApiResource):
@web_ns.response(200, "Success", web_ns.models[AudioToTextResponse.__name__])
def post(self, app_model: App, end_user: EndUser):
"""Convert audio to text"""
file = request.files["file"]
file = request.files.get("file")
try:
response = AudioService.transcript_asr(
@@ -625,7 +625,11 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
message=message_snapshot,
user=user,
stream=stream,
draft_var_saver_factory=self._get_draft_var_saver_factory(invoke_from, account=user),
draft_var_saver_factory=self._get_draft_var_saver_factory(
invoke_from,
account=user,
tenant_id=application_generate_entity.app_config.tenant_id,
),
)
return AdvancedChatAppGenerateResponseConverter.convert(response=response, invoke_from=invoke_from)
@@ -540,6 +540,9 @@ class AgentAppGenerator(MessageBasedAppGenerator):
base_url=dify_config.AGENT_BACKEND_BASE_URL,
use_fake=dify_config.AGENT_BACKEND_USE_FAKE,
fake_scenario=dify_config.AGENT_BACKEND_FAKE_SCENARIO,
stream_read_timeout_seconds=dify_config.AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS,
stream_max_reconnects=dify_config.AGENT_BACKEND_STREAM_MAX_RECONNECTS,
stream_run_timeout_seconds=dify_config.AGENT_BACKEND_RUN_TIMEOUT_SECONDS,
),
event_adapter=AgentBackendRunEventAdapter(),
session_store=AgentAppRuntimeSessionStore(),
+51 -35
View File
@@ -941,48 +941,64 @@ class AgentAppRunner:
if pending_text:
persist_answer_text(pending_text)
for public_event in self._agent_backend_client.stream_events(run_id):
if queue_manager.is_stopped():
flush_pending_agent_message_text()
self._cancel_run(run_id)
raise GenerateTaskStoppedError()
for internal_event in self._event_adapter.adapt(public_event):
try:
public_events = self._agent_backend_client.stream_events(
run_id,
should_stop=queue_manager.is_stopped,
)
for public_event in public_events:
if queue_manager.is_stopped():
flush_pending_agent_message_text()
self._cancel_run(run_id)
raise GenerateTaskStoppedError()
if internal_event.type in (
AgentBackendInternalEventType.RUN_STARTED,
AgentBackendInternalEventType.STREAM_EVENT,
AgentBackendInternalEventType.AGENT_MESSAGE_DELTA,
):
if isinstance(internal_event, AgentBackendAgentMessageDeltaInternalEvent):
debounced_delta = text_delta_debouncer.push(internal_event.delta)
if debounced_delta:
persist_answer_text(debounced_delta)
continue
if isinstance(internal_event, AgentBackendStreamInternalEvent):
for internal_event in self._event_adapter.adapt(public_event):
if queue_manager.is_stopped():
flush_pending_agent_message_text()
try:
process_recorder.handle_stream_event(internal_event)
except Exception:
db.session.rollback()
logger.warning(
"Failed to persist Agent App process event: run_id=%s message_id=%s event_kind=%s",
run_id,
message_id,
internal_event.event_kind,
exc_info=True,
)
self._cancel_run(run_id)
raise GenerateTaskStoppedError()
if internal_event.type in (
AgentBackendInternalEventType.RUN_STARTED,
AgentBackendInternalEventType.STREAM_EVENT,
AgentBackendInternalEventType.AGENT_MESSAGE_DELTA,
):
if isinstance(internal_event, AgentBackendAgentMessageDeltaInternalEvent):
debounced_delta = text_delta_debouncer.push(internal_event.delta)
if debounced_delta:
persist_answer_text(debounced_delta)
continue
if isinstance(internal_event, AgentBackendStreamInternalEvent):
flush_pending_agent_message_text()
try:
process_recorder.handle_stream_event(internal_event)
except Exception:
db.session.rollback()
logger.warning(
"Failed to persist Agent App process event: run_id=%s message_id=%s event_kind=%s",
run_id,
message_id,
internal_event.event_kind,
exc_info=True,
)
continue
continue
continue
flush_pending_agent_message_text()
terminal = internal_event
break
if terminal is not None:
break
flush_pending_agent_message_text()
terminal = internal_event
break
if terminal is not None:
break
except GenerateTaskStoppedError:
raise
except Exception as error:
flush_pending_agent_message_text()
self._cancel_run(run_id)
if queue_manager.is_stopped():
raise GenerateTaskStoppedError() from error
raise
flush_pending_agent_message_text()
if queue_manager.is_stopped():
self._cancel_run(run_id)
raise GenerateTaskStoppedError()
return terminal, process_recorder
def _cancel_run(self, run_id: str) -> None:
+10 -1
View File
@@ -32,6 +32,7 @@ class _DebuggerDraftVariableSaver:
self,
*,
account: Account,
tenant_id: str,
app_id: str,
node_id: str,
node_type: NodeType,
@@ -39,6 +40,7 @@ class _DebuggerDraftVariableSaver:
enclosing_node_id: str | None = None,
) -> None:
self._account = account
self._tenant_id = tenant_id
self._app_id = app_id
self._node_id = node_id
self._node_type = node_type
@@ -49,6 +51,7 @@ class _DebuggerDraftVariableSaver:
with Session(db.engine) as session, session.begin():
DraftVariableSaverImpl(
session=session,
tenant_id=self._tenant_id,
app_id=self._app_id,
node_id=self._node_id,
node_type=self._node_type,
@@ -287,7 +290,12 @@ class BaseAppGenerator:
@final
@staticmethod
def _get_draft_var_saver_factory(invoke_from: InvokeFrom, account: Account | EndUser) -> DraftVariableSaverFactory:
def _get_draft_var_saver_factory(
invoke_from: InvokeFrom,
account: Account | EndUser,
*,
tenant_id: str,
) -> DraftVariableSaverFactory:
if invoke_from == InvokeFrom.DEBUGGER:
assert isinstance(account, Account)
@@ -300,6 +308,7 @@ class BaseAppGenerator:
) -> DraftVariableSaver:
return _DebuggerDraftVariableSaver(
account=account,
tenant_id=tenant_id,
app_id=app_id,
node_id=node_id,
node_type=node_type,
+31 -4
View File
@@ -21,6 +21,7 @@ from core.app.entities.queue_entities import (
WorkflowQueueMessage,
)
from extensions.ext_redis import redis_client
from graphon.graph_engine.manager import GraphEngineManager
from graphon.runtime import GraphRuntimeState
logger = logging.getLogger(__name__)
@@ -51,6 +52,9 @@ class AppQueueManager(ABC):
self._graph_runtime_state: GraphRuntimeState | None = None
self._stopped_cache: TTLCache[tuple, bool] = TTLCache(maxsize=1, ttl=1)
self._cache_lock = threading.Lock()
self._execution_terminal = threading.Event()
self._abort_sent = threading.Event()
self._lifecycle_lock = threading.Lock()
def listen(self):
"""
@@ -59,7 +63,7 @@ class AppQueueManager(ABC):
"""
# wait for APP_MAX_EXECUTION_TIME seconds to stop listen
listen_timeout = dify_config.APP_MAX_EXECUTION_TIME
start_time = time.time()
start_time = time.monotonic()
last_ping_time: int | float = 0
try:
while True:
@@ -72,8 +76,14 @@ class AppQueueManager(ABC):
except queue.Empty:
continue
finally:
elapsed_time = time.time() - start_time
if elapsed_time >= listen_timeout or self._is_stopped():
elapsed_time = time.monotonic() - start_time
timed_out = elapsed_time >= listen_timeout
manually_stopped = self._is_stopped()
if not self._execution_terminal.is_set() and (timed_out or manually_stopped):
reason = (
f"App execution exceeded {listen_timeout} seconds" if timed_out else "App task was stopped"
)
self._abort_execution(reason)
# publish two messages to make sure the client can receive the stop signal
# and stop listening after the stop signal processed
self.publish(
@@ -84,16 +94,33 @@ class AppQueueManager(ABC):
self.publish(QueuePingEvent(), PublishFrom.TASK_PIPELINE)
last_ping_time = elapsed_time // 10
finally:
if not self._execution_terminal.is_set():
self._abort_execution("Client response stream closed before app execution completed")
self._graph_runtime_state = None # Release reference once consumers finish or close the generator.
def stop_listen(self):
def stop_listen(self, *, execution_terminal: bool = False):
"""
Stop listen to queue
:return:
"""
if execution_terminal:
self._execution_terminal.set()
self._clear_task_belong_cache()
self._q.put(None)
def _abort_execution(self, reason: str) -> None:
"""Propagate response timeout/disconnect to legacy and GraphEngine runners."""
with self._lifecycle_lock:
if self._execution_terminal.is_set() or self._abort_sent.is_set():
return
self._abort_sent.set()
try:
self.set_stop_flag_no_user_check(self._task_id)
GraphEngineManager(redis_client).send_stop_command(self._task_id, reason=reason)
except Exception:
logger.exception("Failed to abort app execution for task %s", self._task_id)
def _clear_task_belong_cache(self) -> None:
"""
Remove the task belong cache key once listening is finished.
@@ -45,7 +45,7 @@ class MessageBasedAppQueueManager(AppQueueManager):
if isinstance(
event, QueueStopEvent | QueueErrorEvent | QueueMessageEndEvent | QueueAdvancedChatMessageEndEvent
):
self.stop_listen()
self.stop_listen(execution_terminal=True)
if pub_from == PublishFrom.APPLICATION_MANAGER and self._is_stopped():
if self._app_mode == AppMode.ADVANCED_CHAT.value:
@@ -349,6 +349,7 @@ class PipelineGenerator(BaseAppGenerator):
draft_var_saver_factory = self._get_draft_var_saver_factory(
invoke_from,
user,
tenant_id=pipeline.tenant_id,
)
# return response or stream generator
response = self._handle_response(
@@ -42,7 +42,7 @@ class PipelineQueueManager(AppQueueManager):
| QueueWorkflowFailedEvent
| QueueWorkflowPartialSuccessEvent,
):
self.stop_listen()
self.stop_listen(execution_terminal=True)
if pub_from == PublishFrom.APPLICATION_MANAGER and self._is_stopped():
raise GenerateTaskStoppedError()
+5 -1
View File
@@ -399,7 +399,11 @@ class WorkflowAppGenerator(BaseAppGenerator):
worker_thread.start()
draft_var_saver_factory = self._get_draft_var_saver_factory(invoke_from, user)
draft_var_saver_factory = self._get_draft_var_saver_factory(
invoke_from,
user,
tenant_id=app_model.tenant_id,
)
# return response or stream generator
response = self._handle_response(
@@ -41,4 +41,4 @@ class WorkflowAppQueueManager(AppQueueManager):
| QueueWorkflowFailedEvent
| QueueWorkflowPartialSuccessEvent,
):
self.stop_listen()
self.stop_listen(execution_terminal=True)
+7 -1
View File
@@ -26,6 +26,7 @@ from core.app.entities.queue_entities import (
QueueNodeSucceededEvent,
QueueReasoningChunkEvent,
QueueRetrieverResourcesEvent,
QueueStopEvent,
QueueTextChunkEvent,
QueueWorkflowFailedEvent,
QueueWorkflowPartialSuccessEvent,
@@ -424,7 +425,12 @@ class WorkflowBasedAppRunner:
QueueWorkflowFailedEvent(error=event.error, exceptions_count=event.exceptions_count)
)
case GraphRunAbortedEvent():
self._publish_event(QueueWorkflowFailedEvent(error=event.reason or "Unknown error", exceptions_count=0))
self._publish_event(
QueueStopEvent(
stopped_by=QueueStopEvent.StopBy.USER_MANUAL,
reason=event.reason or "Workflow execution aborted",
)
)
case GraphRunPausedEvent():
runtime_state = workflow_entry.graph_engine.graph_runtime_state
paused_nodes = runtime_state.get_paused_nodes()
+4
View File
@@ -500,11 +500,15 @@ class QueueStopEvent(AppQueueEvent):
event: QueueEvent = QueueEvent.STOP
stopped_by: StopBy
reason: str | None = None
def get_stop_reason(self) -> str:
"""
To stop reason
"""
if self.reason:
return self.reason
reason_mapping = {
QueueStopEvent.StopBy.USER_MANUAL: "Stopped by user.",
QueueStopEvent.StopBy.ANNOTATION_REPLY: "Stopped by annotation reply.",
+91 -2
View File
@@ -47,6 +47,24 @@ class MaxRetriesExceededError(ValueError):
pass
class ResponseLimitError(ValueError):
"""Base error for responses that cannot be safely bounded."""
pass
class ResponseTooLargeError(ResponseLimitError):
"""Raised when an identity response exceeds the configured byte limit."""
pass
class UnsupportedResponseEncodingError(ResponseLimitError):
"""Raised when response encoding prevents safe decoded-size enforcement."""
pass
request_error = httpx.RequestError
max_retries_exceeded_error = MaxRetriesExceededError
@@ -142,7 +160,31 @@ def _inject_trace_headers(headers: Headers | None) -> Headers:
return headers
def make_request(method: str, url: str, max_retries: int = SSRF_DEFAULT_MAX_RETRIES, **kwargs: Any) -> httpx.Response:
def make_request(
method: str,
url: str,
max_retries: int = SSRF_DEFAULT_MAX_RETRIES,
stream_response: bool = False,
**kwargs: Any,
) -> httpx.Response:
"""Send one SSRF-protected request with optional streaming.
Args:
method: HTTP method sent through the configured SSRF client.
url: Absolute request URL.
max_retries: Number of retry attempts after the initial request.
stream_response: Return an open streaming response that the caller must close.
**kwargs: Additional keyword arguments forwarded to ``httpx.Client``.
Returns:
A buffered response, or an open response when ``stream_response`` is true.
Raises:
ToolSSRFError: The configured SSRF proxy rejects the destination.
MaxRetriesExceededError: All configured request attempts fail.
httpx.RequestError: A request fails while retries are disabled.
ValueError: The SSL verification option or request headers are invalid.
"""
# Convert requests-style allow_redirects to httpx-style follow_redirects
if "allow_redirects" in kwargs:
allow_redirects = kwargs.pop("allow_redirects")
@@ -175,6 +217,11 @@ def make_request(method: str, url: str, max_retries: int = SSRF_DEFAULT_MAX_RETR
# When using a forward proxy, httpx may override the Host header based on the URL.
# We extract and preserve any explicitly set Host header to support virtual hosting.
user_provided_host = _get_user_provided_host_header(headers)
send_kwargs: dict[str, Any] = {}
if "auth" in kwargs:
send_kwargs["auth"] = kwargs.pop("auth")
if "follow_redirects" in kwargs:
send_kwargs["follow_redirects"] = kwargs.pop("follow_redirects")
retries = 0
while retries <= max_retries:
@@ -185,7 +232,11 @@ def make_request(method: str, url: str, max_retries: int = SSRF_DEFAULT_MAX_RETR
if user_provided_host is not None:
headers["host"] = user_provided_host
kwargs["headers"] = headers
response = client.request(method=method, url=url, **kwargs)
request = client.build_request(method=method, url=url, **kwargs)
if stream_response:
response = client.send(request, stream=True, **send_kwargs)
else:
response = client.send(request, **send_kwargs)
# Check for SSRF protection by Squid proxy
if response.status_code in (401, 403):
@@ -195,6 +246,7 @@ def make_request(method: str, url: str, max_retries: int = SSRF_DEFAULT_MAX_RETR
# Squid typically identifies itself in Server or Via headers
if "squid" in server_header or "squid" in via_header:
response.close()
raise ToolSSRFError(
f"Access to '{url}' was blocked by SSRF protection. "
f"The URL may point to a private or local network address. "
@@ -208,6 +260,7 @@ def make_request(method: str, url: str, max_retries: int = SSRF_DEFAULT_MAX_RETR
response.status_code,
url,
)
response.close()
except httpx.RequestError as e:
logger.warning("Request to URL %s failed on attempt %s: %s", url, retries + 1, e)
@@ -220,6 +273,42 @@ def make_request(method: str, url: str, max_retries: int = SSRF_DEFAULT_MAX_RETR
raise MaxRetriesExceededError(f"Reached maximum retries ({max_retries}) for URL {url}")
def buffer_response(response: httpx.Response, *, max_response_bytes: int) -> httpx.Response:
"""Consume one open identity response under a decoded byte limit and close its stream."""
if max_response_bytes <= 0:
raise ValueError("max_response_bytes must be positive")
try:
content_encoding = response.headers.get("content-encoding", "identity").strip().lower()
if content_encoding not in {"", "identity"}:
raise UnsupportedResponseEncodingError(f"content encoding {content_encoding} cannot be safely bounded")
content = bytearray()
for chunk in response.iter_bytes():
if len(content) + len(chunk) > max_response_bytes:
raise ResponseTooLargeError(f"response exceeded {max_response_bytes} bytes")
content.extend(chunk)
decoded_headers = {
name: value
for name, value in response.headers.items()
if name.lower() not in {"content-encoding", "content-length", "transfer-encoding"}
}
try:
request = response.request
except RuntimeError:
request = None
return httpx.Response(
response.status_code,
headers=decoded_headers,
content=bytes(content),
request=request,
extensions=response.extensions,
history=response.history,
default_encoding=response.default_encoding,
)
finally:
response.close()
def get(url: str, max_retries: int = SSRF_DEFAULT_MAX_RETRIES, **kwargs: Any) -> httpx.Response:
return make_request("GET", url, max_retries=max_retries, **kwargs)
-49
View File
@@ -403,55 +403,6 @@ class LLMGenerator:
return ""
return "\n\n".join(sections) + "\n\n"
@classmethod
def classify_workflow_mode(
cls,
tenant_id: str,
instruction: str,
model_config: ModelConfig,
) -> Literal["workflow", "advanced-chat"]:
"""Classify a free-text instruction into a concrete app mode.
One tiny LLM call using the model the user already picked (so no extra
provider setup is needed). Parsed leniently; defaults to
``advanced-chat`` on anything unexpected or any error, so a
``mode="auto"`` request never blocks generation. NEVER raises.
"""
default_mode: Literal["workflow", "advanced-chat"] = "advanced-chat"
try:
model_instance = ModelManager.for_tenant(tenant_id=tenant_id).get_model_instance(
tenant_id=tenant_id,
model_type=ModelType.LLM,
provider=model_config.provider,
model=model_config.name,
)
prompt_messages: list[PromptMessage] = [
UserPromptMessage(
content=(
"Reply with exactly one word: 'workflow' (one-shot automation, no chat) "
"or 'advanced-chat' (conversational multi-turn). "
f"Instruction: {instruction.strip()}"
)
),
]
response: LLMResult = model_instance.invoke_llm(
prompt_messages=prompt_messages,
model_parameters={"max_tokens": 4, "temperature": 0},
stream=False,
)
text = (response.message.get_text_content() or "").strip().lower()
except Exception:
logger.info("Workflow mode classification failed; defaulting to %s", default_mode, exc_info=True)
return default_mode
# Lenient parse: an affirmative "workflow" wins; everything else
# (including a truncated / empty / garbled reply) falls back to the
# conversational default. "advanced-chat" needs no positive match
# because it IS the default.
if "workflow" in text:
return "workflow"
return default_mode
@classmethod
def generate_rule_config(cls, tenant_id: str, args: RuleGeneratePayload):
output_parser = RuleConfigGeneratorOutputParser()
+4 -1
View File
@@ -131,7 +131,10 @@ class WaterCrawlAPIClient(BaseAPIClient):
content_type = response.headers.get("Content-Type", "")
media_type = content_type.split(";", 1)[0].strip().lower()
if media_type == "application/json":
return response.json() or {}
try:
return response.json() or {}
except ValueError as exc:
raise ValueError("Invalid JSON response from WaterCrawl") from exc
if media_type == "application/octet-stream":
return response.content
@@ -1,5 +1,15 @@
"""WaterCrawl domain exceptions.
These exceptions are constructed from upstream HTTP responses, which may be
JSON API errors or plain text/HTML proxy errors. Keep the exception type stable
even when the body is not JSON so callers can handle WaterCrawl failures by
domain type instead of low-level parser errors.
"""
import json
from typing import override
from typing import Any, override
from httpx import Response
class WaterCrawlError(Exception):
@@ -7,11 +17,16 @@ class WaterCrawlError(Exception):
class WaterCrawlBadRequestError(WaterCrawlError):
def __init__(self, response):
def __init__(self, response: Response):
self.status_code = response.status_code
self.response = response
data = response.json()
self.message = data.get("message", "Unknown error occurred")
try:
data: Any = response.json()
except ValueError:
data = {}
if not isinstance(data, dict):
data = {}
self.message = data.get("message") or response.text or "Unknown error occurred"
self.errors = data.get("errors", {})
super().__init__(self.message)
+4 -4
View File
@@ -4,12 +4,12 @@ Workflow generator package.
Generates a Dify workflow graph (nodes, edges, viewport) from a natural-language
instruction. Intended for the cmd+k `/create` slash command's preview/apply flow.
Pipeline (slim, single-shot variant):
Pipeline:
runner.WorkflowGenerator.generate_workflow_graph(...)
planner_prompts: short LLM call high-level node plan
builder_prompts: structured-output LLM call full graph JSON
postprocess: fill defaults, auto-layout viewport, sanity-check edges
planner_prompts: short LLM call node and edge plan
node_builder_prompts: bounded parallel calls semantic node configs
postprocess: assemble wrappers, auto-layout, validate graph
The runner is pure domain logic; ``WorkflowGeneratorService`` (in ``services/``)
owns the model-manager dependency and is what controllers call.
@@ -1,19 +1,4 @@
"""
Builder prompts.
The builder is the second step of the slim plannerbuilder pipeline. It takes
the planner's high-level node list and emits the *full* graph JSON consumed by
``WorkflowService.sync_draft_workflow``.
The builder owns: node configuration (prompts, code, headers, etc.), edge wiring,
handle ids ("source"/"target"), positions, and the viewport. It is the only
prompt that needs to know the concrete shape of each node type keep its
examples accurate or the LLM will invent fields.
"""
import json
from collections.abc import Iterable
from typing import Any
"""Compact semantic configuration references for workflow node builders."""
# Per-node-type configuration cheatsheet.
#
@@ -23,50 +8,9 @@ from typing import Any
# both ``WorkflowService.sync_draft_workflow``'s structural checks and the
# runtime entity validation each node performs when the workflow runs.
#
# The cheatsheet is assembled DYNAMICALLY per request: the planner decides
# which node types the workflow needs, and ``build_node_config_cheatsheet``
# stitches together only the snippets for those types (plus the always-needed
# wrapper / shared-field / edge-handle preamble, and the containers section
# when an iteration / loop is planned). This keeps the builder prompt tight —
# a 3-node summariser no longer carries the schema for 12 unrelated node
# types — and lets each snippet document its FULL schema (e.g. a "file" start
# variable's required ``allowed_file_types``) without bloating every prompt.
#
# The postprocessor in ``runner.py`` fills missing wrapper fields (``type``,
# ``positionAbsolute``, ``width``, ``height``, ``sourcePosition`` /
# ``targetPosition``, edge ``data.sourceType`` / ``data.targetType``), so the
# LLM only needs to emit semantically meaningful fields.
# Always-included preamble: the node/edge wrapper shape and the shared
# ``data`` fields that apply to every node type, plus the "## Per type" header
# the per-type snippets slot under.
_CHEATSHEET_PREAMBLE = """\
## Node wrapper (every node, top-level)
{"id": "node1" (digits + letters only see "Node IDs" below),
"type": "custom", # ReactFlow renderer key. Iteration/loop
# *start* children use special types
# (see Containers below).
"position": {"x": <number>, "y": <number>},
"data": { ... per-type fields ... }}
Children of iteration / loop containers additionally need
``parentId``, ``zIndex: 1002`` and ``extent: "parent"`` see Containers.
## Shared "data" fields (every node)
{"type": "<node-type>", # e.g. "llm", "start", "if-else"
"title": "<short label>",
"desc": "<one-liner>",
"selected": false}
## Per type — additional "data" fields (only the node types in your plan are shown)"""
# node_type → its per-type schema snippet. Keyed by the exact ``node_type``
# string the planner emits so ``build_node_config_cheatsheet`` can look each
# one up directly. Iteration / loop are documented in the Containers section
# (they are subgraphs, not leaf nodes) rather than here.
# Each snippet mirrors the production node default closely enough for one
# model call to emit only meaningful ``data`` fields. The runner owns wrappers,
# topology, container metadata, layout, and validation.
_NODE_SNIPPETS: dict[str, str] = {
"start": """\
- start:
@@ -248,484 +192,30 @@ _NODE_SNIPPETS: dict[str, str] = {
Enable only the sub-features you need; ``conditions`` reuse the if-else
condition shape (key / comparison_operator / value). Outputs: ``result``
(the processed array), ``first_record``, ``last_record``.""",
"assigner": """\
- assigner (write to an existing conversation / loop variable):
{"version": "2",
"items": [{"variable_selector": ["<target-node>", "<target-var>"],
"input_type": "variable",
"operation": "over-write",
"value": ["<source-node>", "<source-var>"]}]}
``input_type`` is "variable" (value is a selector) or "constant".
Operations: over-write | clear | append | extend | set | += | -= | *= |
/= | remove-first | remove-last.""",
"human-input": """\
- human-input (pause for a person; use webapp delivery by default):
{"delivery_methods": [{"id": "webapp", "type": "webapp", "enabled": true}],
"form_content": "<short review / approval instructions>",
"inputs": [{"type": "paragraph", "output_variable_name": "comment",
"default": {"type": "constant", "selector": [], "value": ""}}],
"user_actions": [{"id": "approve", "title": "Approve",
"button_style": "primary"}],
"timeout": 3, "timeout_unit": "day"}
Each ``inputs[].output_variable_name`` is an output variable. Outgoing
edges use the matching user-action id as ``sourceHandle``.""",
}
# Pulled into the cheatsheet only when an iteration / loop appears in the plan.
_CONTAINERS_SECTION = """\
## Containers — iteration / loop
These are SUBGRAPH nodes. To use one you MUST emit, in order:
1. The container node itself, e.g. for iteration:
id: "nodeK"
type: "custom"
data: {"type": "iteration",
"title": "<label>",
"desc": "",
"selected": false,
"start_node_id": "nodeKstart",
"iterator_selector": ["<src>", "<list-var>"],
"output_selector": ["<inner-last-node>", "<out-var>"],
"is_parallel": false,
"parallel_nums": 10,
"error_handle_mode": "terminated",
"flatten_output": true}
width: 808
height: 204
zIndex: 1
For loop, swap "iteration" "loop" and use:
data: {"type": "loop", "title": "...", "desc": "",
"selected": false, "start_node_id": "nodeKstart",
"break_conditions": [], "loop_count": 10,
"logical_operator": "and"}
2. The auto-start child (one per container):
id: "nodeKstart"
type: "custom-iteration-start" # loop → "custom-loop-start"
parentId: "nodeK"
extent: "parent"
draggable: false
selectable: false
zIndex: 1002
position: {"x": 60, "y": 78} # relative to parent
data: {"type": "iteration-start", # loop → "loop-start"
"title": "", "desc": "",
"isInIteration": true, # loop → "isInLoop": true
"selected": false}
3. Each inner-pipeline node (any node type, follows normal data rules) MUST add:
parentId: "nodeK"
extent: "parent"
zIndex: 1002
position: {x, y} # relative to parent
data: {..., "isInIteration": true, # loop → "isInLoop": true
"iteration_id": "nodeK"} # loop → "loop_id"
4. Edges INSIDE a container must add to ``data``:
"isInIteration": true # loop → "isInLoop": true
"iteration_id": "nodeK" # loop → "loop_id"
and use ``zIndex: 1002``. Edges OUTSIDE containers use the default
``isInIteration: false`` / ``isInLoop: false``.
5. The container's incoming/outgoing edges connect to the container's id
(``nodeK``), NOT to inner nodes. The first inner edge connects from
``nodeKstart``."""
# Always-included trailer: edge handle conventions for every graph.
_EDGE_HANDLES_SECTION = """\
## Edge handles
- Most nodes: sourceHandle "source", targetHandle "target".
- if-else cases: sourceHandle is the case_id ("true" / "false" / ...).
- question-classifier: sourceHandle is the class_id ("1" / "2" / ...).
- iteration-start / sourceHandle "source"; the edge from the *start node
loop-start: is what kicks off the first inner step."""
# Container node types are described in ``_CONTAINERS_SECTION`` rather than as
# leaf snippets; their presence in a plan pulls that section in.
_CONTAINER_NODE_TYPES = frozenset({"iteration", "loop"})
def build_node_config_cheatsheet(node_types: Iterable[str] | None = None) -> str:
"""
Assemble the builder cheatsheet for exactly the node types in the plan.
``node_types`` is the set of ``node_type`` strings the planner chose. We
emit the always-on preamble (wrapper / shared fields), then only the
per-type snippets for the requested types (``start`` is always included
every graph has one), the Containers section when an iteration / loop is
planned, and the edge-handles trailer. Unknown / unrecognised type strings
are ignored (the runtime / structural validator catches genuinely bogus
types).
``None`` returns the FULL cheatsheet (every snippet + containers) used to
build the static back-compat constants below and as a safe fallback.
"""
if node_types is None:
requested: set[str] = set(_NODE_SNIPPETS) | set(_CONTAINER_NODE_TYPES)
else:
requested = {str(t).strip() for t in node_types if str(t).strip()}
requested.add("start") # every workflow has exactly one start node
parts: list[str] = [_CHEATSHEET_PREAMBLE]
# Iterate _NODE_SNIPPETS (not ``requested``) to keep a stable, readable order.
parts.extend(snippet for node_type, snippet in _NODE_SNIPPETS.items() if node_type in requested)
if requested & _CONTAINER_NODE_TYPES:
parts.append(_CONTAINERS_SECTION)
parts.append(_EDGE_HANDLES_SECTION)
return "\n\n".join(parts) + "\n"
# Full cheatsheet (all node types) — retained as a module constant so callers
# and tests that want the complete reference can import it directly. The
# dynamic per-request prompt is built by ``get_builder_system_prompt``.
NODE_CONFIG_CHEATSHEET = build_node_config_cheatsheet()
_BASE_SYSTEM_PROMPT_HEAD = """You are a Dify workflow builder.
You are given:
1. A user instruction (what the workflow should do).
2. A node plan from the planner (which nodes to use, in execution order).
Your job: emit a complete Dify workflow graph as JSON. The graph will be written
directly into a Studio draft, so it must be syntactically valid and structurally
correct.
# Hard rules
1. The output is a single JSON object no prose, no Markdown, no code fences.
2. NODE IDs MUST USE ONLY ALPHANUMERICS + UNDERSCORES never hyphens.
Dify's run-time placeholder regex (see ``variable_pool.VARIABLE_PATTERN``)
is ``\\{\\{#([a-zA-Z0-9_]{1,50}(?:\\.[a-zA-Z_][a-zA-Z0-9_]{0,29}){1,10})#\\}\\}``,
so any placeholder pointing at a hyphenated id (e.g. ``{{#node-1.text#}}``)
silently fails to match at run time and the literal string survives into
the prompt the user then sees ``{{#node-1.text#}}`` in their output.
Use the EXACT ids from the plan, formatted as ``node1``, ``node2``, ... in
plan order. Edge ``source`` / ``target`` must reference these ids.
3. Every node has top-level fields: id, type, position, data.
- "type" is always "custom" (ReactFlow node renderer).
- "data.type" is the actual node type ("llm", "start", etc.).
4. Every edge has top-level fields: id, source, target, type, sourceHandle, targetHandle.
- "type" is always "custom".
- "sourceHandle"/"targetHandle" follow the cheatsheet (default: "source"/"target").
- Edge id format: "<source>-<sourceHandle>-<target>-<targetHandle>".
5. Use the model from the planner context for ALL "llm" / "question-classifier" /
"parameter-extractor" nodes (provider, name, mode, completion_params).
6. Reference upstream outputs with the literal placeholder syntax
``{{#<node-id>.<output-var>#}}`` — that's DOUBLE curly braces with ``#``
markers inside (matching Dify's runtime placeholder regex
``\\{\\{#[^#]+#\\}\\}``). NEVER emit single-brace ``{#…#}`` — Dify will
not interpolate it, so the LLM at run time would see the literal
placeholder string in its prompt and echo it back as output. Use
``["<node-id>", "<output-var>"]`` for ``value_selector`` /
``query_variable_selector`` / etc.
7. The "start" node owns input variables; downstream nodes reference them as
``["<start-node-id>", "<var-name>"]`` for selectors or
``{{#<start-node-id>.<var-name>#}}`` inside prompt strings.
8. NEVER emit "code" or "http-request" nodes if a tool from the "Available tools"
section below covers the same task replace them with a "tool" node referencing
the exact provider/tool identifier from the catalogue. "code" / "http-request"
are last-resort escape hatches for arbitrary transformations and APIs that no
installed tool can express.
9. EVERY variable reference MUST resolve to a real, declared variable on the
source node never invent a variable name. Specifically:
- ``{{#<node-id>.<var>#}}`` inside a prompt / ``answer`` / ``template-transform``
template (DOUBLE braces single ``{#…#}`` is NOT a Dify placeholder
and will NOT be substituted), AND ``["<node-id>", "<var>"]`` inside a
``value_selector`` /
``query_variable_selector`` / ``iterator_selector`` / ``output_selector`` /
``tool_parameters[*].value`` (when ``type: "variable"``), MUST point at a
value that the source node actually exposes:
* ``start`` one of the ``data.variables[*].variable`` entries you
declared on the start node. Add an entry if you need a new input.
* ``llm`` ``text`` (the default LLM output) or, when structured
output is enabled, a key from its schema.
* ``code`` a key in ``data.outputs``.
* ``knowledge-retrieval`` ``result`` (the standard array output).
* ``parameter-extractor`` one of the ``data.parameters[*].name``.
* ``document-extractor`` ``text`` (extracted file text; an array of
strings when ``is_array_file`` is true).
* ``variable-aggregator`` ``output``.
* ``list-operator`` ``result`` (array), ``first_record``,
``last_record``.
* ``tool`` any parameter declared by the tool the run time
validates these, so you can name them freely, but pick from the
documented provider/tool.
If the planner's "Start inputs" list (see user prompt) is non-empty,
copy each entry verbatim into ``start.data.variables`` so the
downstream references resolve.
- In Advanced-Chat mode you may also reference ``sys.query`` and
``sys.files`` without declaring them. In selector fields, spell these as
exactly ``["sys", "query"]`` and ``["sys", "files"]`` never as a
one-item array such as ``["sys,query"]`` or ``["sys.query"]``.
10. MULTIPLE KNOWLEDGE-RETRIEVAL INPUTS TO ONE LLM require a template fan-in.
``context.variable_selector accepts only one selector`` and therefore
cannot carry two retrieval outputs. When an LLM must synthesize two or
more retrieval results:
- Run the retrieval nodes as parallel siblings from the same query input.
- Add one ``template-transform`` node after them. Give it one variable per
retrieval, such as ``value_selector: ["node2", "result"]`` and
``value_selector: ["node3", "result"]``, and render every source's
content into one labelled text output.
- Add an edge from EACH retrieval node to the template, then one edge from
the template to the LLM. The LLM must NOT receive direct retrieval edges.
- Enable the LLM's context, set ``context.variable_selector`` to the
template's ``["<template-node-id>", "output"]``, and put
``{{#context#}}`` in its prompt.
- Do not use ``variable-aggregator`` for this: it selects the first value
produced by mutually exclusive branches; it does not concatenate two
retrieval results that both ran.
"""
_BASE_SYSTEM_PROMPT_TAIL = """\
# Layout
- Place nodes left-to-right with x=80 + 320 * index, y=280.
- Viewport: {"x": 0, "y": 0, "zoom": 0.7}.
"""
_BASE_SYSTEM_PROMPT_FOOTER = """
# Output schema
{
"nodes": [...],
"edges": [...],
"viewport": {"x": 0, "y": 0, "zoom": 0.7}
}
"""
_WORKFLOW_MODE_RULES = """# Mode-specific rules — Workflow
- The graph MUST start with exactly one "start" node and end with exactly one "end" node.
- Do NOT use "answer" nodes (those are for Advanced Chat only).
- The "end" node's outputs[].value_selector must point at a real upstream output.
"""
_ADVANCED_CHAT_MODE_RULES = """# Mode-specific rules — Advanced Chat (Chatflow)
- The graph MUST start with exactly one "start" node and end with exactly one "answer" node.
- Do NOT use "end" nodes (those are for plain Workflow apps).
- The "start" node should expose "sys.query" / "sys.files" automatically; user-defined
variables go in start.data.variables.
- The "answer" node's "answer" field references upstream outputs as
{{#<node-id>.<var>#}} and is what the user sees in chat.
"""
def _assemble_builder_system_prompt(mode: str, node_types: Iterable[str] | None) -> str:
"""Stitch the builder system prompt for ``mode`` around a cheatsheet built
for ``node_types`` (``None`` full cheatsheet)."""
mode_rules = _ADVANCED_CHAT_MODE_RULES if mode == "advanced-chat" else _WORKFLOW_MODE_RULES
return (
_BASE_SYSTEM_PROMPT_HEAD
+ mode_rules
+ _BASE_SYSTEM_PROMPT_TAIL
+ build_node_config_cheatsheet(node_types)
+ _BASE_SYSTEM_PROMPT_FOOTER
)
# Static full-cheatsheet prompts — the back-compat default returned by
# ``get_builder_system_prompt`` when the caller doesn't pin a node-type set.
BUILDER_SYSTEM_PROMPT_WORKFLOW = _assemble_builder_system_prompt("workflow", None)
BUILDER_SYSTEM_PROMPT_ADVANCED_CHAT = _assemble_builder_system_prompt("advanced-chat", None)
BUILDER_USER_PROMPT = """# User instruction
{instruction}
{ideal_output_section}\
{existing_graph_section}\
# Selected model (use for all LLM-based nodes)
provider={provider}, name={name}, mode={mode_label}
{tool_catalogue_section}\
{start_inputs_section}\
# Node plan (from planner — use these labels and node_types in this order)
{plan_block}
Now emit the complete workflow graph JSON.
"""
# Node wrapper fields that carry no meaning the builder needs: pure canvas /
# selection state, plus geometry the runner's postprocess recomputes anyway.
# Stripping them out of the refine prompt cuts its size roughly in half on
# hand-edited graphs — fewer tokens in, and (because the builder echoes
# untouched nodes verbatim) far fewer tokens out, which is where the latency
# lives.
_PRUNED_NODE_KEYS = frozenset(
{
"positionAbsolute",
"sourcePosition",
"targetPosition",
"selected",
"dragging",
"measured",
}
)
# Additionally pruned from TOP-LEVEL nodes only: the layered auto-layout
# recomputes their position and size defaults, so the builder never needs to
# reproduce them. Container children keep ``position`` (relative to the
# parent, which we cannot recompute) and containers keep ``width`` /
# ``height`` (their canvas size is real config, not a default).
_PRUNED_TOP_LEVEL_NODE_KEYS = _PRUNED_NODE_KEYS | {"position", "width", "height"}
_CONTAINER_DATA_TYPES = frozenset({"iteration", "loop"})
# Edge fields the builder must echo; everything else (ids, zIndex,
# sourceType / targetType, isInIteration / isInLoop markers) is recomputed
# by the runner's postprocess from the node topology.
_KEPT_EDGE_KEYS = ("source", "target", "sourceHandle", "targetHandle")
def compact_graph_for_builder(current_graph: dict) -> dict:
"""
Strip canvas noise out of a draft graph before prompt injection.
Keeps everything semantically meaningful ids, wrapper ``type``,
``parentId``, the full ``data`` config, child positions, container
sizes and drops geometry / selection state the postprocess pass
recomputes. The builder echoes untouched nodes verbatim, so every byte
removed here is removed twice (prompt AND completion).
"""
nodes_out: list[dict] = []
for node in current_graph.get("nodes") or []:
if not isinstance(node, dict):
continue
is_child = bool(node.get("parentId"))
is_container = isinstance(node.get("data"), dict) and node["data"].get("type") in _CONTAINER_DATA_TYPES
pruned = _PRUNED_NODE_KEYS if (is_child or is_container) else _PRUNED_TOP_LEVEL_NODE_KEYS
compact = {k: v for k, v in node.items() if k not in pruned}
if is_container:
# Container position is still recomputed by the layout pass.
compact.pop("position", None)
nodes_out.append(compact)
edges_out = [
{k: edge[k] for k in _KEPT_EDGE_KEYS if k in edge}
for edge in (current_graph.get("edges") or [])
if isinstance(edge, dict)
]
return {"nodes": nodes_out, "edges": edges_out}
def format_builder_existing_graph_section(current_graph: dict | None) -> str:
"""
Refine mode: give the builder the existing graph JSON so it can keep
every node and edge the user's change does not touch byte-for-byte — same
ids, same config, same prompt templates. Without the full config the
builder would regenerate untouched nodes from scratch and silently drop
the user's hand-tuned settings. Canvas-only fields are stripped first
(see ``compact_graph_for_builder``) they're recomputed in postprocess,
so carrying them only slows the call down.
Returns an empty string in create mode (no ``current_graph``); the builder
then behaves exactly as before, constructing the graph purely from the
planner's node plan.
"""
if not current_graph:
return ""
graph_json = json.dumps(compact_graph_for_builder(current_graph), ensure_ascii=False, separators=(",", ":"))
return (
"# Existing graph to refine (JSON)\n\n"
"You are REFINING this existing graph, NOT building from scratch. Apply "
"ONLY the change the user instruction describes. Every node and edge the "
"change does not affect MUST be preserved verbatim — keep the same node "
"ids, the same `data` config, and the same prompt templates. The node "
"plan below is the target node set after your change; use the existing "
"graph as the source of truth for the config of nodes that carry over.\n\n"
f"```json\n{graph_json}\n```\n\n"
)
def format_start_inputs_section(start_inputs: list[dict[str, Any]]) -> str:
"""
Surface the planner's ``start_inputs`` list to the builder so it can
populate ``start.data.variables`` with the exact set of inputs every
downstream variable reference will need. Empty list empty section,
because the builder may then declare no input variables (e.g. an
Advanced-Chat workflow that only consumes ``sys.query``).
"""
if not start_inputs:
return ""
lines = ["# Start inputs (copy each entry verbatim into start.data.variables)"]
lines.append("")
for inp in start_inputs:
variable = str(inp.get("variable") or "").strip()
label = str(inp.get("label") or "").strip()
type_ = str(inp.get("type") or "paragraph").strip()
if not variable:
continue
lines.append(f"- variable={variable!r} label={label!r} type={type_!r}")
lines.append("")
return "\n".join(lines) + "\n"
def format_builder_tool_catalogue_section(catalogue_text: str) -> str:
"""
Builder-facing catalogue block. The builder needs the same identifiers
the planner saw, plus a stern reminder that ``tool`` nodes MUST set
``provider_id`` / ``provider_name`` / ``tool_name`` to entries that
actually exist in this list hallucinated tools fail at draft sync.
"""
if not catalogue_text.strip():
return ""
return (
"# Available tools (use these exact provider/tool identifiers — "
"for each 'tool' node, set provider_id and provider_name to the "
"provider portion and tool_name to the tool portion)\n\n"
f"{catalogue_text}\n\n"
)
def format_plan_block(plan_nodes: list[dict[str, Any]]) -> str:
"""
Render the planner output as a numbered list the builder can quote.
Node IDs use no separator (``node1``, ``node2``, ...) because Dify's
run-time placeholder regex requires ``[a-zA-Z0-9_]`` in the node-id
slot a hyphenated id like ``node-1`` would silently fail to match
at run time and the literal ``{{#node-1.var#}}`` survives into the
LLM prompt.
For container children (planner emitted a ``"parent": "<label>"`` key),
we resolve the parent label to its ``nodeN`` id and surface it on the
same line so the builder knows to set ``parentId`` and the
``isInIteration`` / ``isInLoop`` markers on inner nodes.
"""
# First pass: label → node-id so we can resolve "parent" hints.
label_to_id: dict[str, str] = {}
for idx, node in enumerate(plan_nodes, start=1):
label = str(node.get("label") or "")
if label and label not in label_to_id:
label_to_id[label] = f"node{idx}"
lines = []
for idx, node in enumerate(plan_nodes, start=1):
node_id = f"node{idx}"
label = node.get("label", "")
node_type = node.get("node_type", "")
purpose = node.get("purpose", "")
parent_label = str(node.get("parent") or "")
parent_clause = ""
if parent_label:
parent_id = label_to_id.get(parent_label, "")
if parent_id:
parent_clause = f" parent={parent_id}"
else:
parent_clause = f" parent={parent_label!r}"
lines.append(f"{idx}. id={node_id} type={node_type} label={label!r}{parent_clause}\n purpose: {purpose}")
return "\n".join(lines)
def get_builder_system_prompt(mode: str, node_types: Iterable[str] | None = None) -> str:
"""
Build the builder system prompt for ``mode``, with a cheatsheet scoped to
``node_types`` (the planner's chosen node types).
When ``node_types`` is ``None`` we return the cached full-cheatsheet
constant (back-compat default). When the runner passes the plan's node-type
set we assemble a fresh prompt carrying only the relevant per-type schemas,
so the builder isn't handed config for node types the workflow never uses.
"""
if node_types is None:
return BUILDER_SYSTEM_PROMPT_ADVANCED_CHAT if mode == "advanced-chat" else BUILDER_SYSTEM_PROMPT_WORKFLOW
return _assemble_builder_system_prompt(mode, node_types)
def get_node_config_snippet(node_type: str) -> str:
"""Return the semantic config reference for one leaf node type."""
return _NODE_SNIPPETS.get(node_type, "")
@@ -0,0 +1,138 @@
"""Compact prompts for parallel, per-node workflow configuration.
Each call produces only the semantic ``data`` fields for one planned node.
Canvas wrappers, shared labels, topology, layout, and edge defaults are owned
by ``WorkflowGenerator`` so completion length scales with node configuration
rather than with the full ReactFlow graph.
"""
import json
from typing import Any
from core.workflow.generator.prompts.builder_prompts import get_node_config_snippet
_CONTAINER_CONFIG_SNIPPETS = {
"iteration": """- iteration:
{"iterator_selector": ["<src>", "<list-var>"],
"output_selector": ["<last-child>", "<out-var>"],
"is_parallel": false, "parallel_nums": 10,
"error_handle_mode": "terminated", "flatten_output": true}
The runner supplies start_node_id, child wrappers, and the synthetic start node.""",
"loop": """- loop:
{"break_conditions": [{"id": "c1",
"variable_selector": ["<child>", "<var>"],
"comparison_operator": "is",
"value": "<value>"}],
"loop_count": 10, "logical_operator": "and"}
The runner supplies start_node_id, child wrappers, and the synthetic start node.""",
}
_NODE_BUILDER_HEAD = """You configure exactly ONE node in a Dify workflow.
Return one JSON object with exactly this shape: {"config": {...}}.
``config`` contains only node-type-specific ``data`` fields. Do NOT repeat id,
type, title, desc, selected, position, wrapper fields, edges, or viewport.
Rules:
- Use only ids from the supplied normalized plan.
- Placeholder strings use ``{{#node_id.variable#}}``; selector fields use
``["node_id", "variable"]``. Never invent an upstream output.
- Use the selected model verbatim for llm, question-classifier, and
parameter-extractor nodes.
- Keep prompts/code concise but complete for the user's requested behavior.
- Emit strict JSON only: no prose, Markdown, comments, or trailing commas.
# Target node schema
"""
NODE_BUILDER_USER_PROMPT = """# Target node
id={node_id}, type={node_type}, label={label!r}
purpose={purpose}
# User instruction
{instruction}
{ideal_output_section}{mode_section}{model_section}{tool_catalogue_section}{start_inputs_section}{existing_config_section}\
# Normalized plan and topology
{plan_json}
Return {{"config": {{...}}}} for target node {node_id} now.
"""
def get_node_builder_system_prompt(node_type: str) -> str:
"""Build a one-node prompt containing only that node's semantic schema."""
snippet = _CONTAINER_CONFIG_SNIPPETS.get(node_type) or get_node_config_snippet(node_type)
return _NODE_BUILDER_HEAD + (snippet or f"- {node_type}: emit the minimum valid config fields.")
def format_parallel_plan(
plan_nodes: list[dict[str, Any]],
plan_edges: list[dict[str, Any]],
start_inputs: list[dict[str, Any]] | None = None,
) -> str:
"""Serialize the shared plan compactly so every node call has graph context.
``start_inputs`` rides along so downstream builders reference the declared
``{{#<start-id>.<variable>#}}`` names instead of guessing them from prose
a guessed name gets auto-injected as a spurious form input later.
"""
payload: dict[str, Any] = {"nodes": plan_nodes, "edges": plan_edges}
if start_inputs:
payload["start_inputs"] = start_inputs
return json.dumps(payload, ensure_ascii=False, separators=(",", ":"))
def format_mode_section(mode: str) -> str:
"""Tell each builder which app mode it is configuring for.
Matters most in advanced-chat, where ``sys.query`` / ``sys.files`` are the
sanctioned way to reference the user's message — without this the model
invents start-node variables that postprocess then materializes as
spurious form inputs.
"""
if mode == "advanced-chat":
return (
"# App mode\n\n"
"advanced-chat: the user's chat message is available as sys.query and uploaded files "
'as sys.files — placeholder {{#sys.query#}}, selector ["sys", "query"]. Reference them '
"directly; do NOT invent start-node variables for the chat message.\n\n"
)
return (
"# App mode\n\n"
"workflow: there are NO automatic system variables; reference user input only through "
"the start node's declared variables.\n\n"
)
def format_start_inputs_section(start_inputs: list[dict[str, Any]]) -> str:
"""Render planner-declared inputs for the start-node builder only."""
if not start_inputs:
return ""
lines = ["# Start inputs (copy each entry verbatim into start.data.variables)", ""]
for input_ in start_inputs:
variable = str(input_.get("variable") or "").strip()
if not variable:
continue
label = str(input_.get("label") or "").strip()
type_ = str(input_.get("type") or "paragraph").strip()
lines.append(f"- variable={variable!r} label={label!r} type={type_!r}")
lines.append("")
return "\n".join(lines) + "\n"
def format_tool_catalogue_section(catalogue_text: str) -> str:
"""Render exact tool identifiers for a tool-node builder only."""
if not catalogue_text.strip():
return ""
return (
"# Available tools (use these exact provider/tool identifiers — "
"set provider_id and provider_name to the provider portion and "
"tool_name to the tool portion)\n\n"
f"{catalogue_text}\n\n"
)
@@ -1,13 +1,14 @@
"""
Planner prompts.
The planner is the lightweight first step in the slim plannerbuilder pipeline.
The planner is the lightweight first step in the slim plannernode-builders pipeline.
It receives the user's natural-language instruction and emits a high-level
node plan in JSON. The builder later turns that plan into the final graph.
node and edge plan in JSON. Node builders later produce configs that the runner
assembles into the final graph.
We keep the planner deliberately short the heavy lifting (config schemas,
edge wiring, default values) belongs in the builder. The planner only commits
to the *which-node-types* decision so the builder gets a tight scaffold.
default values) belongs in the builders. The planner commits to the minimum
topology and node types so every builder gets a tight scaffold.
"""
PLANNER_SYSTEM_PROMPT = """You are a Dify workflow planner.
@@ -39,6 +40,8 @@ minimum set of Dify workflow nodes needed to fulfil it, in execution order.
mutually-exclusive paths before "end" / "answer".
- "list-operator" filter / sort / slice an array variable (e.g. the items
fed into or produced by an "iteration").
- "assigner" update an existing conversation or loop variable.
- "human-input" pause for a person to review, approve, or enter data.
# Rules
@@ -97,22 +100,45 @@ minimum set of Dify workflow nodes needed to fulfil it, in execution order.
variables are automatic downstream nodes may reference them without
a ``start_inputs`` entry. In Workflow mode there is NO automatic
variable; everything the user supplies must be in ``start_inputs``.
11. Output strictly the JSON object no prose, no Markdown, no code fences.
11. Give every node a unique runtime-safe ``id`` using only letters, digits,
and underscores. In create mode use ``node1``, ``node2``, ... in node-list
order. In refine mode preserve the existing id for every retained node.
12. Emit the target graph's edges in ``edges``. Each edge is
``{"source": "<id>", "target": "<id>"}``; add ``source_handle`` only
for branch nodes: if-else case id, question-classifier class id, or
human-input action id. Container children reference the container id in
their ``parent`` field; do not emit the synthetic iteration/loop start node.
13. In refine mode add ``action`` to every retained target node:
``"keep"`` when its data config is unchanged, ``"update"`` when the user
asked to change its config, and ``"add"`` for a new node. Removed nodes are
omitted. Edge-only rewiring does not require changing a node's action.
14. Output strictly the JSON object no prose, no Markdown, no code fences.
15. Echo the app mode in the ``mode`` output field exactly "workflow" or
"advanced-chat". When the ``# Mode`` section says auto, YOU decide:
"workflow" for one-shot automations (run once with form inputs, return a
result), "advanced-chat" for conversational multi-turn assistants. The
terminal node must match the chosen mode (rule 2): "end" for workflow,
"answer" for advanced-chat.
# Output schema
{
"title": "<≤ 40-char title of the workflow>",
"description": "<one-sentence summary>",
"mode": "workflow | advanced-chat",
"app_name": "<≤ 30-char product-style name, e.g. 'URL Summarizer'>",
"icon": "<single emoji that captures the workflow's purpose, e.g. '📰'>",
"start_inputs": [
{"variable": "url", "label": "URL", "type": "text-input"}
],
"nodes": [
{"label": "Start", "node_type": "start", "purpose": "..."},
{"label": "Summarize", "node_type": "llm", "purpose": "..."},
{"label": "End", "node_type": "end", "purpose": "..."}
{"id": "node1", "label": "Start", "node_type": "start", "purpose": "..."},
{"id": "node2", "label": "Summarize", "node_type": "llm", "purpose": "..."},
{"id": "node3", "label": "End", "node_type": "end", "purpose": "..."}
],
"edges": [
{"source": "node1", "target": "node2"},
{"source": "node2", "target": "node3"}
]
}
"""
@@ -140,7 +166,8 @@ def format_existing_graph_section(current_graph: dict | None) -> str:
We pass only ids / node-types / titles + edge endpoints here the planner
decides *which nodes* exist, so it needs the shape, not the per-node config.
The builder gets the full graph JSON to preserve untouched node config.
Node builders receive only the config of a node marked ``update``;
configs marked ``keep`` are reused directly.
"""
if not current_graph:
return ""
@@ -156,7 +183,13 @@ def format_existing_graph_section(current_graph: dict | None) -> str:
for edge in edges:
if not isinstance(edge, dict):
continue
edge_lines.append(f"- {edge.get('source', '')} -> {edge.get('target', '')}")
# Branch wiring (if-else case ids, classifier class ids, human-input
# action ids) lives in ``sourceHandle``. The planner is the only
# source of edges for the rebuilt graph, so the real handle must be
# surfaced here or refine silently rewires branches.
handle = str(edge.get("sourceHandle") or "")
handle_suffix = f" (source_handle={handle!r})" if handle and handle != "source" else ""
edge_lines.append(f"- {edge.get('source', '')} -> {edge.get('target', '')}{handle_suffix}")
nodes_block = "\n".join(node_lines) or "(none)"
edges_block = "\n".join(edge_lines) or "(none)"
return (
@@ -166,7 +199,9 @@ def format_existing_graph_section(current_graph: dict | None) -> str:
"node list to reflect that change while keeping everything the "
"instruction does not mention — preserve existing nodes, their order, "
"and their labels wherever the change leaves them untouched. Only add, "
"remove, or rename nodes the requested change actually requires.\n\n"
"remove, or rename nodes the requested change actually requires. "
"For every retained edge, copy its source_handle verbatim from the "
"list below — branch wiring must survive the refine unchanged.\n\n"
f"Current nodes:\n{nodes_block}\n\n"
f"Current edges:\n{edges_block}\n\n"
)
+468 -119
View File
@@ -1,14 +1,14 @@
"""
Workflow generator runner.
Slim plannerbuilder pipeline. Pure domain logic; the model instance is
Slim plannerparallel-node-builder pipeline. Pure domain logic; the model instance is
injected by ``WorkflowGeneratorService`` so this module stays cleanly
separated from the infrastructure layer.
Pipeline:
1. PLANNER short LLM call producing a high-level node list.
2. BUILDER structured-output LLM call producing the full graph JSON.
2. BUILDERS bounded concurrent LLM calls producing compact node configs.
3. POSTPROC fill safe defaults, lay nodes out left-to-right, dedupe
edge ids, and run a final structural sanity check.
@@ -20,7 +20,8 @@ Intentionally NOT here (deferred to a future iteration):
- Tool / model catalogue filtering
If quality regresses below product threshold we add those back; for now the
single planner+builder pair shipped behind cmd+k `/create` is enough.
planner and bounded parallel node builders shipped behind cmd+k `/create` are
enough.
"""
import json
@@ -28,17 +29,22 @@ import logging
import re
import time
from collections.abc import Iterator
from concurrent.futures import ThreadPoolExecutor, as_completed
from copy import deepcopy
from typing import Any, ClassVar, cast
import json_repair
from core.workflow.generator.prompts.builder_prompts import (
BUILDER_USER_PROMPT,
format_builder_existing_graph_section,
format_builder_tool_catalogue_section,
format_plan_block,
from configs import dify_config
from core.workflow.generator.prompts.node_builder_prompts import (
NODE_BUILDER_USER_PROMPT,
format_mode_section,
format_parallel_plan,
format_start_inputs_section,
get_builder_system_prompt,
get_node_builder_system_prompt,
)
from core.workflow.generator.prompts.node_builder_prompts import (
format_tool_catalogue_section as format_node_tool_catalogue_section,
)
from core.workflow.generator.prompts.planner_prompts import (
PLANNER_SYSTEM_PROMPT,
@@ -55,6 +61,7 @@ from core.workflow.generator.types import (
WorkflowGenerateErrorDict,
WorkflowGenerateResultDict,
WorkflowGenerationMode,
WorkflowGenerationModeRequest,
)
from graphon.enums import BuiltinNodeTypes
from graphon.model_runtime.entities.llm_entities import LLMResult
@@ -96,11 +103,31 @@ _DEFAULT_FILE_UPLOAD_METHODS = ("local_file", "remote_url")
# Token ceiling for the planner call when the caller didn't pin one. The plan
# is a short JSON node list (a handful of nodes with labels/purposes), so this
# is generous headroom while still bounding a runaway response. The builder is
# left on the caller's budget — it emits the full graph and genuinely needs it.
# is generous headroom while still bounding a runaway response. Builder calls
# keep the caller's budget so complex node configs are not truncated.
_PLANNER_DEFAULT_MAX_TOKENS = 4096
# Per-node calls trade a larger request count for a shorter critical path.
# The cap comes from ``WORKFLOW_GENERATOR_NODE_BUILDER_MAX_WORKERS`` (default 6,
# enough to run the planner's recommended 36-node plans as a single wave);
# provider rate-limit bursts are absorbed by ``_invoke_with_retry``'s bounded
# backoff, and operators can dial the env var back down if their provider is
# stricter. Read at call time so tests (and live config reloads) can adjust it
# without re-importing.
def _node_builder_max_workers() -> int:
return dify_config.WORKFLOW_GENERATOR_NODE_BUILDER_MAX_WORKERS
_MODEL_NODE_TYPES = frozenset(
{
BuiltinNodeTypes.LLM,
BuiltinNodeTypes.QUESTION_CLASSIFIER,
BuiltinNodeTypes.PARAMETER_EXTRACTOR,
}
)
# Appended as a trailing user message on the SECOND (and only) attempt when
# the first response wasn't parseable as JSON. Keep this terse — the model
# already has its full instructions in the original system message; this is
@@ -110,6 +137,12 @@ _JSON_RETRY_HINT = (
"Do not include any prose, markdown code fences, comments, or trailing commas."
)
_PLANNER_SCHEMA_RETRY_HINT = (
"Your plan did not match the required topology schema. Return the complete plan again with "
"a unique non-empty id on every node and a non-empty edges array whose source and target "
"reference those ids. Return ONLY the JSON object."
)
# Provider hiccups we retry: a dropped connection, a 5xx, or a rate-limit are
# all transient — the same request usually succeeds moments later. We do NOT
@@ -189,14 +222,58 @@ def _result_with_errors(
def _with_mode(result: WorkflowGenerateResultDict, mode: WorkflowGenerationMode) -> WorkflowGenerateResultDict:
"""Stamp the resolved concrete ``mode`` onto a result envelope.
``mode="auto"`` requests are resolved to a concrete mode before planning;
echoing it back lets the frontend pick the right app type to create. It's
present for explicit modes too so the response shape stays uniform.
``mode="auto"`` requests are resolved to a concrete mode from the planner
output; echoing it back lets the frontend pick the right app type to
create. It's present for explicit modes too so the response shape stays
uniform.
"""
result["mode"] = mode
return result
def _fallback_mode(mode: WorkflowGenerationModeRequest) -> WorkflowGenerationMode:
"""Concrete mode for envelopes emitted before the planner resolved one.
``auto`` maps to the conversational default the same never-fail fallback
the old standalone classifier used so ``result.mode`` never leaks the
``auto`` sentinel to the frontend.
"""
return "advanced-chat" if mode == "auto" else mode
def _planner_prompt_mode(mode: WorkflowGenerationModeRequest) -> str:
"""Mode string interpolated into the planner user prompt.
For ``auto`` the value is self-describing so the planner knows the choice
is delegated to it (system-prompt rule 15).
"""
return "auto (choose workflow or advanced-chat)" if mode == "auto" else mode
def _resolve_generation_mode(
requested: WorkflowGenerationModeRequest, plan: PlannerResultDict
) -> WorkflowGenerationMode:
"""Resolve the request mode into the concrete generation mode.
An explicit request always wins a contradictory planner ``mode`` field is
ignored. For ``auto``: trust the planner's echoed ``mode``, else infer from
the plan's terminal node type (the structural source of truth the graph is
validated against), else fall back to the conversational default. Lenient
on purpose a bad ``mode`` value must never fail the plan.
"""
if requested != "auto":
return requested
planner_mode = str(plan.get("mode") or "").strip().lower()
if planner_mode in ("workflow", "advanced-chat"):
return cast(WorkflowGenerationMode, planner_mode)
node_types = {str(node.get("node_type") or "") for node in plan.get("nodes") or [] if isinstance(node, dict)}
if BuiltinNodeTypes.ANSWER in node_types:
return "advanced-chat"
if BuiltinNodeTypes.END in node_types:
return "workflow"
return "advanced-chat"
def _build_plan_event(
*,
plan: PlannerResultDict,
@@ -255,7 +332,7 @@ class WorkflowGenerator:
provider: str,
model_name: str,
model_mode: str,
mode: WorkflowGenerationMode,
mode: WorkflowGenerationModeRequest,
instruction: str,
ideal_output: str = "",
tool_catalogue_text: str = "",
@@ -263,20 +340,26 @@ class WorkflowGenerator:
current_graph: dict[str, Any] | None = None,
) -> WorkflowGenerateResultDict:
"""
Run planner builder postprocess and return a graph payload.
Run planner node builders postprocess and return a graph payload.
``mode`` accepts the ``"auto"`` sentinel the planner then chooses the
concrete mode itself (echoed in its ``mode`` output field) so no extra
classification call is needed; the resolution is stamped onto the
result envelope.
``current_graph`` switches the pipeline from create mode to REFINE
mode: the existing draft graph is injected into both the planner
(compact node/edge summary) and the builder (full JSON) so the LLM
amends the graph the user is editing instead of inventing a new one.
``None`` (the default) is plain create-from-scratch behaviour.
mode: the existing draft graph is summarized for the planner. Node
builders receive only the config of the node they update, while configs
marked ``keep`` are reused without an LLM call. ``None`` (the default)
is plain create-from-scratch behaviour.
``tool_catalogue_text`` is the formatted list of installed tools for
the calling tenant (see ``tool_catalogue.build_tool_catalogue`` /
``format_tool_catalogue``). It's injected into both the planner and
builder prompts so the LLM can pick concrete ``provider/tool``
identifiers instead of inventing names; an empty string skips the
section entirely (useful for unit tests).
identifiers instead of inventing names; node builders receive it
only for tool nodes. An empty string skips the section entirely (useful
for unit tests).
``installed_tools`` is the structural sibling a set of
``(provider_name, tool_name)`` pairs the validator consults to reject
@@ -316,7 +399,7 @@ class WorkflowGenerator:
# The event generator always emits exactly one result envelope; this
# fallback only guards against a future refactor that forgets to.
if result is None:
result = _with_mode(_empty_result(), mode)
result = _with_mode(_empty_result(), _fallback_mode(mode))
return result
@classmethod
@@ -328,7 +411,7 @@ class WorkflowGenerator:
provider: str,
model_name: str,
model_mode: str,
mode: WorkflowGenerationMode,
mode: WorkflowGenerationModeRequest,
instruction: str,
ideal_output: str = "",
tool_catalogue_text: str = "",
@@ -368,7 +451,7 @@ class WorkflowGenerator:
provider: str,
model_name: str,
model_mode: str,
mode: WorkflowGenerationMode,
mode: WorkflowGenerationModeRequest,
instruction: str,
ideal_output: str = "",
tool_catalogue_text: str = "",
@@ -376,7 +459,7 @@ class WorkflowGenerator:
current_graph: dict[str, Any] | None = None,
) -> Iterator[tuple[str, dict[str, Any]]]:
"""
Drive planner builder postprocess and yield generation events.
Drive planner node builders postprocess and yield generation events.
Shared core for both ``generate_workflow_graph`` (keeps only the final
``result``) and ``generate_workflow_graph_stream`` (streams every
@@ -402,11 +485,16 @@ class WorkflowGenerator:
),
)
if plan_err is not None:
yield "result", cast(dict[str, Any], _with_mode(_result_with_errors(_empty_result(), [plan_err]), mode))
failed = _with_mode(_result_with_errors(_empty_result(), [plan_err]), _fallback_mode(mode))
yield "result", cast(dict[str, Any], failed)
return
# The lambda return is non-None when no error fired — narrow it for type-checkers.
plan = cast(PlannerResultDict, plan)
# ``auto`` requests resolve here — the planner echoed its mode choice
# (or we infer it from the plan's terminal node). Explicit modes pass
# through unchanged. Everything downstream uses the concrete mode.
resolved_mode = _resolve_generation_mode(mode, plan)
plan_nodes: list[dict[str, Any]] = cast(list[dict[str, Any]], plan.get("nodes", []))
if not plan_nodes:
empty_plan = _with_mode(
@@ -414,15 +502,12 @@ class WorkflowGenerator:
_empty_result(),
[_err(WorkflowGenerateErrorCode.EMPTY_PLAN, "Planner returned no nodes")],
),
mode,
resolved_mode,
)
yield "result", cast(dict[str, Any], empty_plan)
return
# A single LLM cannot select multiple retrieval outputs as context.
# Make the required template fan-in explicit in the plan so the
# builder receives its schema and assigns stable sequential ids.
cls._insert_multi_retrieval_template_plan(plan_nodes)
plan_edges = [cast(dict[str, Any], edge) for edge in (plan.get("edges") or []) if isinstance(edge, dict)]
# Planner-supplied user-input declarations. The builder uses these to
# populate ``start.data.variables`` so downstream ``{#start.<var>#}``
@@ -436,34 +521,48 @@ class WorkflowGenerator:
# First event the stream sees: the high-level plan, before the slower
# builder call. Non-streaming callers ignore it.
yield "plan", _build_plan_event(plan=plan, plan_nodes=plan_nodes, start_inputs=start_inputs, mode=mode)
yield "plan", _build_plan_event(plan=plan, plan_nodes=plan_nodes, start_inputs=start_inputs, mode=resolved_mode)
# ── 2. BUILDER ────────────────────────────────────────────────────
graph, build_err = cls._run_stage(
stage="Builder",
failure_fallback_message="Failed to build workflow graph",
run=lambda: cls._run_builder(
builder_started_at = time.monotonic()
def build_graph() -> GraphDict:
return cls._run_parallel_node_builders(
model_instance=model_instance,
model_parameters=model_parameters,
provider=provider,
model_name=model_name,
model_mode=model_mode,
mode=mode,
mode=resolved_mode,
instruction=instruction,
ideal_output=ideal_output,
plan_nodes=plan_nodes,
plan_edges=plan_edges,
tool_catalogue_text=tool_catalogue_text,
start_inputs=start_inputs,
current_graph=current_graph,
),
)
graph, build_err = cls._run_stage(
stage="Builder",
failure_fallback_message="Failed to build workflow graph",
run=build_graph,
)
logger.info(
"Workflow generator: node builders completed nodes=%s elapsed_ms=%.1f",
len(plan_nodes),
(time.monotonic() - builder_started_at) * 1000,
)
if build_err is not None:
yield "result", cast(dict[str, Any], _with_mode(_result_with_errors(_empty_result(), [build_err]), mode))
yield (
"result",
cast(dict[str, Any], _with_mode(_result_with_errors(_empty_result(), [build_err]), resolved_mode)),
)
return
graph = cast(GraphDict, graph)
# ── 3. POSTPROC + VALIDATE ────────────────────────────────────────
graph = cls._postprocess_graph(graph=graph, mode=mode)
graph = cls._postprocess_graph(graph=graph, mode=resolved_mode)
# ``app_name`` / ``icon`` are planner display metadata; both default
# to "" when the LLM omits them — the FE owns the fallback.
@@ -475,13 +574,13 @@ class WorkflowGenerator:
"error": "",
"errors": [],
}
_with_mode(result, mode)
_with_mode(result, resolved_mode)
# Final structural sanity check — fail closed if start/end shape is
# wrong, container topology is broken, a tool was hallucinated, or a
# variable reference points at a node that won't expose it. We still
# return the partial graph so the caller can debug or salvage it.
structural_errors = cls._validate_structure(graph=graph, mode=mode, installed_tools=installed_tools)
structural_errors = cls._validate_structure(graph=graph, mode=resolved_mode, installed_tools=installed_tools)
if structural_errors:
logger.warning("Workflow generator: structural validation failed: %s", structural_errors)
yield "result", cast(dict[str, Any], _result_with_errors(result, structural_errors))
@@ -622,14 +721,14 @@ class WorkflowGenerator:
*,
model_instance,
model_parameters: dict[str, Any],
mode: WorkflowGenerationMode,
mode: WorkflowGenerationModeRequest,
instruction: str,
ideal_output: str,
tool_catalogue_text: str,
current_graph: dict[str, Any] | None = None,
) -> PlannerResultDict:
user_prompt = PLANNER_USER_PROMPT.format(
mode=mode,
mode=_planner_prompt_mode(mode),
instruction=instruction.strip(),
existing_graph_section=format_existing_graph_section(current_graph),
ideal_output_section=format_ideal_output_section(ideal_output),
@@ -639,59 +738,65 @@ class WorkflowGenerator:
SystemPromptMessage(content=PLANNER_SYSTEM_PROMPT),
UserPromptMessage(content=user_prompt),
]
clamped_parameters = _clamp_for_planner(model_parameters)
parsed = cls._invoke_and_parse_json(
model_instance=model_instance,
messages=messages,
model_parameters=_clamp_for_planner(model_parameters),
model_parameters=clamped_parameters,
stage="Planner",
)
try:
return cls._validate_planner_schema(parsed)
except _StageSchemaError:
logger.info("Workflow generator: planner schema invalid; retrying once")
parsed = cls._invoke_and_parse_json(
model_instance=model_instance,
messages=[*messages, UserPromptMessage(content=_PLANNER_SCHEMA_RETRY_HINT)],
model_parameters=clamped_parameters,
stage="Planner",
)
return cls._validate_planner_schema(parsed)
@staticmethod
def _validate_planner_schema(parsed: dict[str, Any]) -> PlannerResultDict:
"""Require the single planner contract consumed by node builders."""
nodes = parsed.get("nodes")
if not isinstance(nodes, list):
raise _StageSchemaError("Planner", "missing 'nodes' array")
if not nodes:
return cast(PlannerResultDict, parsed)
node_ids: set[str] = set()
for node in nodes:
if not isinstance(node, dict) or "node_type" not in node:
if not isinstance(node, dict) or not node.get("node_type"):
raise _StageSchemaError("Planner", f"malformed node entry: {node!r}")
node_id = node.get("id")
if not isinstance(node_id, str) or not node_id.strip():
raise _StageSchemaError("Planner", f"node missing non-empty id: {node!r}")
if node_id in node_ids:
raise _StageSchemaError("Planner", f"duplicate node id: {node_id!r}")
node_ids.add(node_id)
edges = parsed.get("edges")
if not isinstance(edges, list) or not edges:
raise _StageSchemaError("Planner", "missing non-empty 'edges' array")
for edge in edges:
if not isinstance(edge, dict):
raise _StageSchemaError("Planner", f"malformed edge entry: {edge!r}")
source = edge.get("source")
target = edge.get("target")
if not isinstance(source, str) or not isinstance(target, str):
raise _StageSchemaError("Planner", f"edge missing source or target: {edge!r}")
if source not in node_ids or target not in node_ids:
raise _StageSchemaError("Planner", f"edge references unknown node: {edge!r}")
return cast(PlannerResultDict, parsed)
# ------------------------------------------------------------------
# Plan normalization
# ------------------------------------------------------------------
@staticmethod
def _insert_multi_retrieval_template_plan(plan_nodes: list[dict[str, Any]]) -> None:
"""Insert the unambiguous multi-retrieval template step when omitted.
Multiple LLM nodes make ownership ambiguous, so that case remains in
the planner's hands. With exactly one LLM, every independent retrieval
result can safely fan into one template immediately before that LLM.
"""
node_types = [str(node.get("node_type") or "") for node in plan_nodes]
if node_types.count(BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL) < 2:
return
if node_types.count(BuiltinNodeTypes.LLM) != 1:
return
if BuiltinNodeTypes.TEMPLATE_TRANSFORM in node_types:
return
llm_index = node_types.index(BuiltinNodeTypes.LLM)
retrievals_before_llm = node_types[:llm_index].count(BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL)
if retrievals_before_llm < 2:
return
plan_nodes.insert(
llm_index,
{
"label": "Combine Knowledge",
"node_type": BuiltinNodeTypes.TEMPLATE_TRANSFORM,
"purpose": "Combine every knowledge retrieval result into one labelled context for the LLM.",
},
)
# ------------------------------------------------------------------
# Builder
# ------------------------------------------------------------------
@classmethod
def _run_builder(
def _run_parallel_node_builders(
cls,
*,
model_instance,
@@ -703,53 +808,275 @@ class WorkflowGenerator:
instruction: str,
ideal_output: str,
plan_nodes: list[dict[str, Any]],
plan_edges: list[dict[str, Any]],
tool_catalogue_text: str,
start_inputs: list[dict[str, Any]] | None = None,
current_graph: dict[str, Any] | None = None,
start_inputs: list[dict[str, Any]],
current_graph: dict[str, Any] | None,
) -> GraphDict:
user_prompt = BUILDER_USER_PROMPT.format(
"""Build changed node configs concurrently and expand them into a graph.
Refine plans can mark existing nodes as ``keep``; those nodes bypass
the model entirely and retain their full data config. Every other node
gets one compact call, with at most ``_node_builder_max_workers()``
calls in flight. Any fragment failure aborts the graph, preserving the
generator's existing fail-closed contract.
"""
existing_by_id = {
str(node.get("id")): node
for node in ((current_graph or {}).get("nodes") or [])
if isinstance(node, dict) and node.get("id")
}
existing_edges = [edge for edge in ((current_graph or {}).get("edges") or []) if isinstance(edge, dict)]
nodes_to_build = [
node for node in plan_nodes if not (node.get("action") == "keep" and str(node.get("id")) in existing_by_id)
]
# Shared across every builder call in this request — compute once.
plan_json = format_parallel_plan(plan_nodes, plan_edges, start_inputs)
mode_section = format_mode_section(mode)
configs_by_id: dict[str, dict[str, Any]] = {}
if nodes_to_build:
max_workers = min(_node_builder_max_workers(), len(nodes_to_build))
with ThreadPoolExecutor(max_workers=max_workers, thread_name_prefix="workflow-node-builder") as executor:
futures = {
executor.submit(
cls._run_node_builder,
model_instance=model_instance,
model_parameters=model_parameters,
provider=provider,
model_name=model_name,
model_mode=model_mode,
mode_section=mode_section,
instruction=instruction,
ideal_output=ideal_output,
target_node=node,
plan_json=plan_json,
tool_catalogue_text=tool_catalogue_text,
start_inputs=start_inputs,
existing_node=existing_by_id.get(str(node.get("id"))),
): str(node.get("id"))
for node in nodes_to_build
}
try:
for future in as_completed(futures):
node_id = futures[future]
configs_by_id[node_id] = future.result()
except BaseException:
# Fail fast: one failed fragment aborts the whole graph, so
# queued builder calls would only burn quota and delay the
# error envelope. In-flight calls cannot be interrupted;
# they finish while the pool shuts down.
for pending in futures:
pending.cancel()
raise
return cls._assemble_parallel_graph(
plan_nodes=plan_nodes,
plan_edges=plan_edges,
configs_by_id=configs_by_id,
existing_by_id=existing_by_id,
existing_edges=existing_edges,
)
@classmethod
def _run_node_builder(
cls,
*,
model_instance,
model_parameters: dict[str, Any],
provider: str,
model_name: str,
model_mode: str,
mode_section: str,
instruction: str,
ideal_output: str,
target_node: dict[str, Any],
plan_json: str,
tool_catalogue_text: str,
start_inputs: list[dict[str, Any]],
existing_node: dict[str, Any] | None,
) -> dict[str, Any]:
"""Generate only the semantic config for one normalized plan node."""
node_id = str(target_node.get("id") or "")
node_type = str(target_node.get("node_type") or "")
model_section = ""
if node_type in _MODEL_NODE_TYPES:
model_section = (
f"# Selected model (copy verbatim)\n\nprovider={provider}, name={model_name}, mode={model_mode}\n\n"
)
existing_config_section = ""
if existing_node:
existing_data = existing_node.get("data") if isinstance(existing_node.get("data"), dict) else {}
existing_config_section = (
"# Existing config to preserve unless the instruction changes it\n\n"
f"{json.dumps(existing_data, ensure_ascii=False, separators=(',', ':'))}\n\n"
)
user_prompt = NODE_BUILDER_USER_PROMPT.format(
node_id=node_id,
node_type=node_type,
label=str(target_node.get("label") or ""),
purpose=str(target_node.get("purpose") or ""),
instruction=instruction.strip(),
ideal_output_section=format_ideal_output_section(ideal_output),
existing_graph_section=format_builder_existing_graph_section(current_graph),
provider=provider,
name=model_name,
mode_label=model_mode,
plan_block=format_plan_block(plan_nodes),
tool_catalogue_section=format_builder_tool_catalogue_section(tool_catalogue_text),
start_inputs_section=format_start_inputs_section(start_inputs or []),
mode_section=mode_section,
model_section=model_section,
tool_catalogue_section=(
format_node_tool_catalogue_section(tool_catalogue_text) if node_type == BuiltinNodeTypes.TOOL else ""
),
start_inputs_section=(
format_start_inputs_section(start_inputs) if node_type == BuiltinNodeTypes.START else ""
),
existing_config_section=existing_config_section,
plan_json=plan_json,
)
# Scope the builder cheatsheet to exactly the node types the planner
# chose, so the prompt carries each type's FULL schema (e.g. a file
# start variable's required ``allowed_file_types``) without dragging in
# config for unrelated node types.
plan_node_types = {
str(node.get("node_type") or "").strip() for node in plan_nodes if str(node.get("node_type") or "").strip()
}
messages = [
SystemPromptMessage(content=get_builder_system_prompt(mode, plan_node_types)),
UserPromptMessage(content=user_prompt),
]
parsed = cls._invoke_and_parse_json(
model_instance=model_instance,
messages=messages,
messages=[
SystemPromptMessage(content=get_node_builder_system_prompt(node_type)),
UserPromptMessage(content=user_prompt),
],
model_parameters=model_parameters,
stage="Builder",
stage=f"Builder {node_id}",
)
config = parsed.get("config")
if not isinstance(config, dict):
raise _StageSchemaError(f"Builder {node_id}", "missing 'config' object")
return cast(dict[str, Any], config)
nodes = parsed.get("nodes")
edges = parsed.get("edges")
if not isinstance(nodes, list) or not isinstance(edges, list):
raise _StageSchemaError("Builder", "graph missing 'nodes' or 'edges' arrays")
@classmethod
def _assemble_parallel_graph(
cls,
*,
plan_nodes: list[dict[str, Any]],
plan_edges: list[dict[str, Any]],
configs_by_id: dict[str, dict[str, Any]],
existing_by_id: dict[str, dict[str, Any]],
existing_edges: list[dict[str, Any]] | None = None,
) -> GraphDict:
"""Expand compact node configs and planner topology into graph JSON.
viewport = parsed.get("viewport") or _DEFAULT_VIEWPORT
return cast(
GraphDict,
{
"nodes": nodes,
"edges": edges,
"viewport": viewport,
},
)
``existing_edges`` (refine only) preserves wiring the planner cannot
express: the synthetic ``<container>start`` entry edge keeps its
existing target instead of being re-pointed at whichever child the
planner happened to list first.
"""
label_to_id = {
str(node.get("label")): str(node.get("id")) for node in plan_nodes if node.get("label") and node.get("id")
}
type_by_id = {str(node.get("id")): str(node.get("node_type") or "") for node in plan_nodes}
children_by_parent: dict[str, list[str]] = {}
nodes: list[dict[str, Any]] = []
for planned in plan_nodes:
node_id = str(planned.get("id") or "")
node_type = str(planned.get("node_type") or "")
existing = existing_by_id.get(node_id)
node: dict[str, Any]
if planned.get("action") == "keep" and existing is not None:
node = deepcopy(existing)
else:
config = dict(configs_by_id.get(node_id) or {})
for shared_key in ("type", "title", "desc", "selected"):
config.pop(shared_key, None)
data: dict[str, Any] = {
"type": node_type,
"title": str(planned.get("label") or node_id),
"desc": str(planned.get("purpose") or ""),
**config,
}
node = deepcopy(existing) if existing is not None else {"id": node_id}
node["id"] = node_id
node["data"] = data
parent_ref = str(planned.get("parent") or "")
parent_id = label_to_id.get(parent_ref, parent_ref)
if not parent_id and str(node.get("parentId") or "") in type_by_id:
# Kept nodes rarely re-state containment — recover the parent
# from the deepcopied wrapper so entry-edge synthesis still
# counts this child.
parent_id = str(node["parentId"])
if parent_id:
child_index = len(children_by_parent.get(parent_id, []))
node["parentId"] = parent_id
node.setdefault("position", {"x": 240 + 260 * child_index, "y": 60})
node.setdefault("data", {})
parent_type = type_by_id.get(parent_id)
if parent_type == BuiltinNodeTypes.ITERATION:
node["data"].setdefault("isInIteration", True)
node["data"].setdefault("iteration_id", parent_id)
elif parent_type == BuiltinNodeTypes.LOOP:
node["data"].setdefault("isInLoop", True)
node["data"].setdefault("loop_id", parent_id)
children_by_parent.setdefault(parent_id, []).append(node_id)
elif node.get("parentId"):
# The container was dropped from the plan: strip the stale
# containment markers so the kept node rejoins the top level
# (and its auto-layout) instead of pointing at a deleted parent.
for wrapper_key in ("parentId", "extent", "zIndex", "position", "positionAbsolute"):
node.pop(wrapper_key, None)
if isinstance(node.get("data"), dict):
for marker_key in ("isInIteration", "iteration_id", "isInLoop", "loop_id"):
node["data"].pop(marker_key, None)
nodes.append(node)
if node_type in cls._CONTAINER_TYPES:
start_id = f"{node_id}start"
node.setdefault("data", {})["start_node_id"] = start_id
node.setdefault("width", 808)
node.setdefault("height", 204)
node.setdefault("zIndex", 1)
is_iteration = node_type == BuiltinNodeTypes.ITERATION
nodes.append(
{
"id": start_id,
"type": "custom-iteration-start" if is_iteration else "custom-loop-start",
"parentId": node_id,
"extent": "parent",
"draggable": False,
"selectable": False,
"zIndex": 1002,
"position": {"x": 60, "y": 78},
"data": {
"type": "iteration-start" if is_iteration else "loop-start",
"title": "",
"desc": "",
"selected": False,
"isInIteration" if is_iteration else "isInLoop": True,
},
}
)
edges: list[dict[str, Any]] = []
for planned_edge in plan_edges:
edge: dict[str, Any] = {
"source": str(planned_edge.get("source") or ""),
"target": str(planned_edge.get("target") or ""),
}
source_handle = planned_edge.get("source_handle") or planned_edge.get("sourceHandle")
target_handle = planned_edge.get("target_handle") or planned_edge.get("targetHandle")
if source_handle:
edge["sourceHandle"] = str(source_handle)
if target_handle:
edge["targetHandle"] = str(target_handle)
edges.append(edge)
# Synthesize each container's entry edge. Refine keeps the existing
# entry target when it is still a child — the planner's node listing
# order says nothing about execution order inside a kept container.
planned_sources = {str(edge.get("source") or "") for edge in edges}
existing_entry_targets = {
str(edge.get("source") or ""): str(edge.get("target") or "") for edge in (existing_edges or [])
}
for parent_id, child_ids in children_by_parent.items():
start_id = f"{parent_id}start"
if not child_ids or start_id in planned_sources:
continue
preferred = existing_entry_targets.get(start_id, "")
entry_target = preferred if preferred in child_ids else child_ids[0]
edges.append({"source": start_id, "target": entry_target})
return cast(GraphDict, {"nodes": nodes, "edges": edges, "viewport": _DEFAULT_VIEWPORT})
# ------------------------------------------------------------------
# Postprocessing
@@ -1225,6 +1552,11 @@ class WorkflowGenerator:
return var == "output"
if node_type == BuiltinNodeTypes.LIST_OPERATOR:
return var in {"result", "first_record", "last_record"}
if node_type == BuiltinNodeTypes.HUMAN_INPUT:
return any(
isinstance(item, dict) and item.get("output_variable_name") == var
for item in (data.get("inputs") or [])
)
# Other node types (if-else, iteration-start, loop-start, ...) don't
# produce outputs of their own.
return False
@@ -1247,6 +1579,15 @@ class WorkflowGenerator:
if isinstance(parameter, dict) and isinstance(parameter.get("name"), str)
]
return parameters[0] if len(parameters) == 1 else None
if node_type == BuiltinNodeTypes.HUMAN_INPUT:
human_outputs: list[str] = []
for item in data.get("inputs") or []:
if not isinstance(item, dict):
continue
output_name = item.get("output_variable_name")
if isinstance(output_name, str):
human_outputs.append(output_name)
return human_outputs[0] if len(human_outputs) == 1 else None
if not isinstance(node_type, str):
return None
single_output_by_type: dict[str, str] = {
@@ -1533,6 +1874,12 @@ class WorkflowGenerator:
for klass in (data.get("classes") or [])
if isinstance(klass, dict) and klass.get("id")
]
elif node_type == BuiltinNodeTypes.HUMAN_INPUT:
branch_handles = [
str(action["id"])
for action in (data.get("user_actions") or [])
if isinstance(action, dict) and action.get("id")
]
else:
continue
@@ -1941,8 +2288,8 @@ class WorkflowGenerator:
"""
Validate iteration / loop topology:
* every container has at least one child whose ``parentId``
points at it;
* every container has at least one executable child whose
``parentId`` points at it;
* every non-container node with a ``parentId`` points at a real
container, not at a non-container node;
* no cycles in the parent chain (a node cannot be its own
@@ -1959,7 +2306,9 @@ class WorkflowGenerator:
if not isinstance(parent, str) or not parent:
continue
if parent in container_ids:
children_by_parent.setdefault(parent, []).append(n.get("id", ""))
node_type = (n.get("data") or {}).get("type")
if node_type not in {"iteration-start", "loop-start"}:
children_by_parent.setdefault(parent, []).append(n.get("id", ""))
elif parent in by_id:
# Parent exists but isn't a container — that's a topology bug.
out.append(
+25 -9
View File
@@ -1,10 +1,10 @@
"""
Typed payloads for workflow generation.
These TypedDicts describe the shape that the planner and builder LLM calls are
required to return after ``json_repair`` parsing. They mirror the runtime
``graph`` shape consumed by ``WorkflowService.sync_draft_workflow`` so the output
can be written straight into a draft workflow without further translation.
These TypedDicts describe the planner payload and the runtime graph assembled
from builder LLM responses after ``json_repair`` parsing. The graph types mirror
the shape consumed by ``WorkflowService.sync_draft_workflow`` so the output can
be written straight into a draft workflow.
"""
from enum import StrEnum
@@ -12,11 +12,10 @@ from typing import Literal, NotRequired, TypedDict
WorkflowGenerationMode = Literal["workflow", "advanced-chat"]
# The mode accepted at the API boundary. ``auto`` is a sentinel that asks the
# service to classify the instruction into a concrete ``WorkflowGenerationMode``
# (one tiny LLM call) BEFORE planning — see
# ``WorkflowGeneratorService._resolve_mode`` and
# ``LLMGenerator.classify_workflow_mode``.
# The mode accepted at the API boundary. ``auto`` is a sentinel that delegates
# the choice to the planner: it echoes a concrete mode in its ``mode`` output
# field (falling back to terminal-node inference, then ``advanced-chat``) —
# see ``runner._resolve_generation_mode``. No extra LLM call is involved.
WorkflowGenerationModeRequest = Literal["workflow", "advanced-chat", "auto"]
@@ -58,9 +57,21 @@ class WorkflowGenerateErrorDict(TypedDict):
class PlannerNodeDict(TypedDict):
"""One node from the planner's high-level plan."""
id: NotRequired[str]
label: str
node_type: str
purpose: str
parent: NotRequired[str]
action: NotRequired[Literal["keep", "update", "add"]]
class PlannerEdgeDict(TypedDict):
"""Compact topology emitted by the planner for parallel node building."""
source: str
target: str
source_handle: NotRequired[str]
target_handle: NotRequired[str]
class PlannerStartInputDict(TypedDict):
@@ -82,10 +93,15 @@ class PlannerResultDict(TypedDict):
title: str
description: str
# Concrete mode the planner chose ("workflow" / "advanced-chat"). Parsed
# leniently — an ``auto`` request infers the mode from the terminal node
# when this is missing or invalid, so a bad value never fails the plan.
mode: NotRequired[str]
app_name: NotRequired[str]
icon: NotRequired[str]
start_inputs: NotRequired[list[PlannerStartInputDict]]
nodes: list[PlannerNodeDict]
edges: NotRequired[list[PlannerEdgeDict]]
class GraphNodePositionDict(TypedDict):
+3
View File
@@ -499,6 +499,9 @@ class DifyNodeFactory(NodeFactory):
base_url=dify_config.AGENT_BACKEND_BASE_URL,
use_fake=dify_config.AGENT_BACKEND_USE_FAKE,
fake_scenario=dify_config.AGENT_BACKEND_FAKE_SCENARIO,
stream_read_timeout_seconds=dify_config.AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS,
stream_max_reconnects=dify_config.AGENT_BACKEND_STREAM_MAX_RECONNECTS,
stream_run_timeout_seconds=dify_config.AGENT_BACKEND_RUN_TIMEOUT_SECONDS,
),
"event_adapter": AgentBackendRunEventAdapter(),
# Agent Files §4.6: reback file outputs from the ToolFile row so
+28 -1
View File
@@ -5,6 +5,7 @@ from collections.abc import Generator, Mapping, Sequence
from typing import TYPE_CHECKING, Any, override
from agenton.compositor import CompositorSessionSnapshot
from dify_agent.protocol import CancelRunRequest
from clients.agent_backend import (
AgentBackendAgentMessageDeltaInternalEvent,
@@ -473,7 +474,10 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
"""
stream_event_count = 0
try:
for public_event in self._agent_backend_client.stream_events(run_id):
for public_event in self._agent_backend_client.stream_events(
run_id,
should_stop=self._is_graph_aborted,
):
stream_event_count += 1
for internal_event in self._event_adapter.adapt(public_event):
if internal_event.type == AgentBackendInternalEventType.RUN_STARTED:
@@ -501,6 +505,7 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
| AgentBackendDeferredToolCallInternalEvent,
):
return internal_event, None
self._cancel_backend_run(run_id, reason="unexpected_event")
return None, self._failure_event(
inputs={},
process_data={},
@@ -509,6 +514,7 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
error_type="agent_backend_stream_error",
)
except AgentBackendError as error:
self._cancel_backend_run(run_id, reason=self._stream_stop_reason())
return None, self._failure_event(
inputs={},
process_data={},
@@ -517,6 +523,7 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
error_type=self._agent_backend_error_type(error),
)
except Exception as error:
self._cancel_backend_run(run_id, reason=self._stream_stop_reason())
return None, self._failure_event(
inputs={},
process_data={},
@@ -525,8 +532,28 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
error_type="agent_backend_stream_error",
)
self._cancel_backend_run(run_id, reason="stream_ended_without_terminal_event")
return None, None
def _is_graph_aborted(self) -> bool:
"""Let Agent SSE consumption observe GraphEngine's cooperative abort state."""
try:
return self.graph_runtime_state.graph_execution.aborted
except (AttributeError, RuntimeError):
return False
def _stream_stop_reason(self) -> str:
return "workflow_graph_aborted" if self._is_graph_aborted() else "event_stream_failed"
def _cancel_backend_run(self, run_id: str, *, reason: str) -> None:
try:
self._agent_backend_client.cancel_run(
run_id,
CancelRunRequest(reason=reason, message="Workflow Agent event consumption stopped"),
)
except Exception:
logger.warning("Failed to cancel Workflow Agent backend run: run_id=%s", run_id, exc_info=True)
@staticmethod
def _record_type_check_metadata(metadata: dict[str, Any], outcome: OutputTypeCheckOutcome) -> None:
# Surface enough detail in metadata for Inspector / debug logs without
+23 -6
View File
@@ -61,16 +61,33 @@ class WorkflowAgentNodeValidator:
@classmethod
def validate_draft_workflow(cls, *, session: Session, workflow: Workflow) -> None:
cls._validate_workflow(session=session, workflow=workflow, require_binding=False)
cls._validate_workflow(
session=session,
workflow=workflow,
require_binding=False,
validate_previous_node_topology=False,
)
@classmethod
def validate_published_workflow(cls, *, session: Session, workflow: Workflow) -> None:
cls._validate_workflow(session=session, workflow=workflow, require_binding=True)
cls._validate_workflow(
session=session,
workflow=workflow,
require_binding=True,
validate_previous_node_topology=True,
)
@classmethod
def _validate_workflow(cls, *, session: Session, workflow: Workflow, require_binding: bool) -> None:
def _validate_workflow(
cls,
*,
session: Session,
workflow: Workflow,
require_binding: bool,
validate_previous_node_topology: bool,
) -> None:
graph = workflow.graph_dict
topology = _WorkflowGraphTopology.from_graph(graph)
topology = _WorkflowGraphTopology.from_graph(graph) if validate_previous_node_topology else None
for node_id, node_data in cls.iter_agent_v2_nodes(graph):
cls._validate_node_schema(node_id=node_id, node_data=node_data)
binding = cls._find_binding(
@@ -185,12 +202,12 @@ class WorkflowAgentNodeValidator:
raise WorkflowAgentNodeValidationError(
f"Workflow Agent node {binding.node_id} has invalid previous node output ref."
)
if topology is None:
continue
if len(selector) < 2:
raise WorkflowAgentNodeValidationError(
f"Workflow Agent node {binding.node_id} has incomplete previous node output ref."
)
if topology is None:
continue
source_node_id = selector[0]
if not topology.has_node(source_node_id):
raise WorkflowAgentNodeValidationError(
+255
View File
@@ -0,0 +1,255 @@
"""Validate Dify Console KnowledgeFS declarations against a pinned OpenAPI document.
The OpenAPI document is exported only during explicit development validation. Runtime declarations live with Dify
product policy; this module validates their transport metadata without generating a complete operation catalog.
"""
from __future__ import annotations
import argparse
import hashlib
import json
import subprocess
import sys
import tempfile
from pathlib import Path
from typing import Any, Literal, TypedDict
API_ROOT = Path(__file__).resolve().parents[1]
if str(API_ROOT) not in sys.path:
sys.path.insert(0, str(API_ROOT))
WORKSPACE_ROOT = API_ROOT.parent
LOCK_PATH = API_ROOT / "knowledge-fs-contract.lock.json"
DEFAULT_REPOSITORY = WORKSPACE_ROOT.parent / "knowledge-fs"
OPENAPI_METHODS = ("delete", "get", "head", "options", "patch", "post", "put", "trace")
PROXY_METHODS = frozenset({"delete", "get", "patch", "post", "put"})
class ContractDeclaration(TypedDict):
"""KnowledgeFS transport contract declared by one Dify Console registry entry."""
operation_id: str
method: str
path: str
required_scope: str | None
response_kind: str
max_response_bytes: int
request_headers: tuple[str, ...]
response_headers: tuple[str, ...]
response_media_types: tuple[str, ...]
type DeclarationField = Literal[
"method",
"path",
"required_scope",
"response_kind",
"max_response_bytes",
"request_headers",
"response_headers",
"response_media_types",
]
DECLARATION_FIELDS: tuple[DeclarationField, ...] = (
"method",
"path",
"required_scope",
"response_kind",
"max_response_bytes",
"request_headers",
"response_headers",
"response_media_types",
)
def main() -> None:
"""Update or verify the pin and validate Console declarations against its OpenAPI document."""
parser = argparse.ArgumentParser()
mode = parser.add_mutually_exclusive_group()
mode.add_argument("--check", action="store_true")
mode.add_argument("--update-lock", action="store_true")
parser.add_argument("--repository", type=Path, default=DEFAULT_REPOSITORY)
args = parser.parse_args()
repository = args.repository.resolve()
lock = json.loads(LOCK_PATH.read_text())
tracked_changes = run("git", "status", "--porcelain", "--untracked-files=no", cwd=repository).strip()
if tracked_changes:
raise RuntimeError("KnowledgeFS checkout must not contain tracked changes during contract export")
commit = run("git", "rev-parse", "HEAD", cwd=repository).strip()
if not args.update_lock and commit != lock["commit"]:
raise RuntimeError(
f"KnowledgeFS checkout mismatch: expected {lock['commit']}, received {commit}. "
"Use the pinned commit or pass --update-lock intentionally."
)
with tempfile.TemporaryDirectory(prefix="dify-knowledge-fs-contract-") as directory:
openapi_path = Path(directory) / "knowledge-fs.openapi.json"
subprocess.run(
["pnpm", "openapi:export", "--", "--output", str(openapi_path)],
cwd=repository,
check=True,
)
openapi_content = openapi_path.read_bytes()
openapi_sha256 = sha256(openapi_content)
if not args.update_lock and openapi_sha256 != lock["openapiSha256"]:
raise RuntimeError(
f"KnowledgeFS OpenAPI hash mismatch: expected {lock['openapiSha256']}, received {openapi_sha256}"
)
document: dict[str, Any] = json.loads(openapi_content)
validate_declarations(document, console_contract_declarations())
if args.update_lock:
LOCK_PATH.write_text(
json.dumps(
{
"commit": commit,
"openapiSha256": openapi_sha256,
"repository": lock["repository"],
},
indent=2,
)
+ "\n"
)
def validate_declarations(document: dict[str, Any], declarations: tuple[ContractDeclaration, ...]) -> None:
"""Validate Dify Console declarations against matching pinned OpenAPI operations."""
operations_by_id: dict[str, list[tuple[str, str, dict[str, Any], dict[str, Any]]]] = {}
for path, path_item in document.get("paths", {}).items():
for method in OPENAPI_METHODS:
operation = path_item.get(method)
if operation is None:
continue
operation_id = operation.get("operationId")
if isinstance(operation_id, str) and operation_id:
operations_by_id.setdefault(operation_id, []).append((method, path, path_item, operation))
declared_ids: set[str] = set()
for declaration in declarations:
operation_id = declaration["operation_id"]
if operation_id in declared_ids:
raise ValueError(f"Dify Console registry has duplicate operationId: {operation_id}")
declared_ids.add(operation_id)
matches = operations_by_id.get(operation_id, [])
if not matches:
raise ValueError(f"KnowledgeFS OpenAPI has no operationId: {operation_id}")
if len(matches) > 1:
raise ValueError(f"KnowledgeFS OpenAPI has duplicate operationId: {operation_id}")
method, path, path_item, operation = matches[0]
if not path.startswith("/"):
raise ValueError(f"KnowledgeFS OpenAPI path must be absolute: {path}")
if method not in PROXY_METHODS:
raise ValueError(f"KnowledgeFS proxy does not support {method.upper()} {path}")
expected: ContractDeclaration = {
"operation_id": operation_id,
"method": method.upper(),
"path": path[1:],
"required_scope": required_scope(operation),
"response_kind": response_kind(operation),
"max_response_bytes": required_max_response_bytes(operation),
"request_headers": request_header_names(path_item, operation),
"response_headers": response_header_names(operation),
"response_media_types": response_media_types(operation),
}
for field in DECLARATION_FIELDS:
expected_value = expected[field]
received_value = declaration[field]
if received_value != expected_value:
raise ValueError(
f"KnowledgeFS operation {operation_id} field {field} drifted: "
f"expected {expected_value!r}, received {received_value!r}"
)
def console_contract_declarations() -> tuple[ContractDeclaration, ...]:
"""Return transport declarations from the runtime Console operation registry."""
from services.knowledge_fs_proxy import KNOWLEDGE_FS_CONSOLE_OPERATIONS
return tuple(
{
"operation_id": operation.operation_id,
"method": operation.method,
"path": operation.path,
"required_scope": operation.required_scope,
"response_kind": operation.response_kind,
"max_response_bytes": operation.max_response_bytes,
"request_headers": operation.request_headers,
"response_headers": operation.response_headers,
"response_media_types": operation.response_media_types,
}
for operation in KNOWLEDGE_FS_CONSOLE_OPERATIONS
)
def response_kind(operation: dict[str, Any]) -> str:
media_types = response_media_types(operation)
if "text/event-stream" in media_types:
return "stream"
if "application/octet-stream" in media_types:
return "binary"
return "buffered"
def response_media_types(operation: dict[str, Any]) -> tuple[str, ...]:
media_types: set[str] = set()
for status, response in operation.get("responses", {}).items():
if status == "2XX" or (len(status) == 3 and status.startswith("2") and status.isdigit()):
media_types.update(response.get("content", {}))
return tuple(sorted(media_types))
def required_scope(operation: dict[str, Any]) -> str | None:
scope = operation.get("x-knowledge-fs-required-scope")
if scope in ("knowledge-spaces:read", "knowledge-spaces:write"):
return scope
if operation.get("security") == []:
return None
raise ValueError(f"KnowledgeFS operation has no supported required scope: {scope}")
def required_max_response_bytes(operation: dict[str, Any]) -> int:
value = operation.get("x-knowledge-fs-max-response-bytes")
if not isinstance(value, int) or isinstance(value, bool) or value <= 0:
raise ValueError(f"KnowledgeFS operation has no valid response byte limit: {value}")
return value
def request_header_names(path_item: dict[str, Any], operation: dict[str, Any]) -> tuple[str, ...]:
names: set[str] = set()
for parameter in [*path_item.get("parameters", []), *operation.get("parameters", [])]:
if "$ref" in parameter:
raise ValueError(f"KnowledgeFS request header references are not supported: {parameter['$ref']}")
if parameter.get("in") == "header":
names.add(parameter["name"].lower())
return tuple(sorted(names))
def response_header_names(operation: dict[str, Any]) -> tuple[str, ...]:
return tuple(
sorted(
{
name.lower()
for response in operation.get("responses", {}).values()
for name in response.get("headers", {})
}
)
)
def sha256(content: bytes) -> str:
return hashlib.sha256(content).hexdigest()
def run(*command: str, cwd: Path) -> str:
return subprocess.run(command, cwd=cwd, check=True, capture_output=True, text=True).stdout
if __name__ == "__main__":
main()
+4
View File
@@ -7,6 +7,9 @@ from .delete_tool_parameters_cache_when_sync_draft_workflow import (
handle as handle_delete_tool_parameters_cache_when_sync_draft_workflow,
)
from .queue_credential_sync_when_tenant_created import handle as handle_queue_credential_sync_when_tenant_created
from .queue_default_plugin_install_when_tenant_created import (
handle as handle_queue_default_plugin_install_when_tenant_created,
)
from .sync_plugin_trigger_when_app_created import handle as handle_sync_plugin_trigger_when_app_created
from .sync_webhook_when_app_created import handle as handle_sync_webhook_when_app_created
from .sync_workflow_schedule_when_app_published import handle as handle_sync_workflow_schedule_when_app_published
@@ -32,6 +35,7 @@ __all__ = [
"handle_create_site_record_when_app_created",
"handle_delete_tool_parameters_cache_when_sync_draft_workflow",
"handle_queue_credential_sync_when_tenant_created",
"handle_queue_default_plugin_install_when_tenant_created",
"handle_sync_plugin_trigger_when_app_created",
"handle_sync_webhook_when_app_created",
"handle_sync_workflow_schedule_when_app_published",
@@ -0,0 +1,22 @@
"""Queue default marketplace plugin installation after tenant creation."""
import logging
from configs import dify_config
from events.tenant_event import tenant_was_created
from tasks.install_default_plugins_task import install_default_plugins_task
logger = logging.getLogger(__name__)
@tenant_was_created.connect
def handle(sender, **kwargs) -> None:
"""Keep tenant creation non-blocking while installing configured plugins asynchronously."""
plugin_ids = dify_config.NEW_USER_DEFAULT_PLUGIN_ID_LIST
if not plugin_ids:
return
try:
install_default_plugins_task.delay(sender.id, plugin_ids)
except Exception:
logger.exception("Failed to queue default plugin installation for tenant %s", sender.id)
+1
View File
@@ -156,6 +156,7 @@ def init_app(app: DifyApp) -> Celery:
"tasks.generate_summary_index_task", # summary index generation
"tasks.regenerate_summary_index_task", # summary index regeneration
"tasks.initialize_created_app_rbac_access_task", # app access initialization
"tasks.install_default_plugins_task", # tenant default plugin installation
"tasks.app_generate.resume_agent_app_task", # ENG-635: Agent v2 chat ask_human resume
"tasks.workflow_run_archive_download_tasks", # workflow-run archive download preparation
]
+5
View File
@@ -0,0 +1,5 @@
{
"commit": "4310e2d582d25e7de58183f27720afab01e123cf",
"openapiSha256": "5827ca930ce38462bfd1b2bef387efbf37eb7ffcaedde4558af2fbaeccbfbc4b",
"repository": "https://github.com/langgenius/knowledge-fs"
}
+33 -12
View File
@@ -16,12 +16,19 @@ from urllib.parse import quote
import boto3
import orjson
from botocore.client import Config
from botocore.exceptions import ClientError
from botocore.exceptions import BotoCoreError, ClientError
from configs import dify_config
logger = logging.getLogger(__name__)
_OBJECT_NOT_FOUND_ERROR_CODES = frozenset({"404", "NoSuchKey", "NotFound"})
def _is_object_not_found_error(error: ClientError) -> bool:
error_code = str(error.response.get("Error", {}).get("Code", ""))
return error_code in _OBJECT_NOT_FOUND_ERROR_CODES
class ArchiveStorageError(Exception):
"""Base exception for archive storage operations."""
@@ -138,10 +145,11 @@ class ArchiveStorage:
response = self.client.get_object(Bucket=self.bucket, Key=key)
return response["Body"].read()
except ClientError as e:
error_code = e.response.get("Error", {}).get("Code")
if error_code == "NoSuchKey":
raise FileNotFoundError(f"Archive object not found: {key}")
raise ArchiveStorageError(f"Failed to download object '{key}': {e}")
if _is_object_not_found_error(e):
raise FileNotFoundError(f"Archive object not found: {key}") from e
raise ArchiveStorageError(f"Failed to download object '{key}': {e}") from e
except BotoCoreError as e:
raise ArchiveStorageError(f"Failed to download object '{key}': {e}") from e
def get_object_stream(self, key: str) -> Generator[bytes, None, None]:
"""
@@ -161,10 +169,11 @@ class ArchiveStorage:
response = self.client.get_object(Bucket=self.bucket, Key=key)
yield from response["Body"].iter_chunks()
except ClientError as e:
error_code = e.response.get("Error", {}).get("Code")
if error_code == "NoSuchKey":
raise FileNotFoundError(f"Archive object not found: {key}")
raise ArchiveStorageError(f"Failed to stream object '{key}': {e}")
if _is_object_not_found_error(e):
raise FileNotFoundError(f"Archive object not found: {key}") from e
raise ArchiveStorageError(f"Failed to stream object '{key}': {e}") from e
except BotoCoreError as e:
raise ArchiveStorageError(f"Failed to stream object '{key}': {e}") from e
def object_exists(self, key: str) -> bool:
"""
@@ -175,12 +184,19 @@ class ArchiveStorage:
Returns:
True if object exists, False otherwise
Raises:
ArchiveStorageError: If storage cannot authoritatively determine object existence
"""
try:
self.client.head_object(Bucket=self.bucket, Key=key)
return True
except ClientError:
return False
except ClientError as e:
if _is_object_not_found_error(e):
return False
raise ArchiveStorageError(f"Failed to check archive object '{key}': {e}") from e
except BotoCoreError as e:
raise ArchiveStorageError(f"Failed to check archive object '{key}': {e}") from e
def delete_object(self, key: str) -> None:
"""
@@ -196,7 +212,12 @@ class ArchiveStorage:
self.client.delete_object(Bucket=self.bucket, Key=key)
logger.debug("Deleted object: %s", key)
except ClientError as e:
raise ArchiveStorageError(f"Failed to delete object '{key}': {e}")
if _is_object_not_found_error(e):
logger.debug("Archive object was already absent: %s", key)
return
raise ArchiveStorageError(f"Failed to delete object '{key}': {e}") from e
except BotoCoreError as e:
raise ArchiveStorageError(f"Failed to delete object '{key}': {e}") from e
def generate_presigned_url(
self,
@@ -0,0 +1,39 @@
"""add step by step tour state
Revision ID: b8c9d0e1f2a3
Revises: 3c9f8e2a1d7b
Create Date: 2026-06-29 12:00:00.000000
"""
import sqlalchemy as sa
from alembic import op
import models
# revision identifiers, used by Alembic.
revision = "b8c9d0e1f2a3"
down_revision = "3c9f8e2a1d7b"
branch_labels = None
depends_on = None
def upgrade():
op.create_table(
"account_step_by_step_tour_states",
sa.Column("id", models.types.StringUUID(), nullable=False),
sa.Column("account_id", models.types.StringUUID(), nullable=False),
sa.Column("first_workspace_id", models.types.StringUUID(), nullable=True),
sa.Column("skipped", sa.Boolean(), server_default=sa.text("false"), nullable=False),
sa.Column("completed_task_ids", models.types.AdjustedJSON(), nullable=False),
sa.Column("manually_enabled_workspace_ids", models.types.AdjustedJSON(), nullable=False),
sa.Column("manually_disabled_workspace_ids", models.types.AdjustedJSON(), nullable=False),
sa.Column("created_at", sa.DateTime(), server_default=sa.text("CURRENT_TIMESTAMP"), nullable=False),
sa.Column("updated_at", sa.DateTime(), server_default=sa.text("CURRENT_TIMESTAMP"), nullable=False),
sa.PrimaryKeyConstraint("id", name="account_step_by_step_tour_state_pkey"),
sa.UniqueConstraint("account_id", name="account_step_by_step_tour_state_account_id_key"),
)
def downgrade():
op.drop_table("account_step_by_step_tour_states")
@@ -0,0 +1,45 @@
"""add workflow-run archive bundle cursor indexes
Revision ID: 3c9f8e2a1d7b
Revises: 7a1c2d9e4b60
Create Date: 2026-07-15 16:00:00.000000
"""
from alembic import op
# revision identifiers, used by Alembic.
revision = "3c9f8e2a1d7b"
down_revision = "7a1c2d9e4b60"
branch_labels = None
depends_on = None
_TABLE_NAME = "workflow_run_archive_bundles"
_INDEXES = (
("workflow_run_archive_bundle_month_id_idx", ("year", "month", "id")),
("workflow_run_archive_bundle_month_shard_id_idx", ("year", "month", "shard", "id")),
)
def _is_postgresql() -> bool:
return op.get_bind().dialect.name == "postgresql"
def upgrade() -> None:
if _is_postgresql():
with op.get_context().autocommit_block():
for index_name, columns in _INDEXES:
op.create_index(index_name, _TABLE_NAME, columns, postgresql_concurrently=True)
return
for index_name, columns in _INDEXES:
op.create_index(index_name, _TABLE_NAME, columns)
def downgrade() -> None:
if _is_postgresql():
with op.get_context().autocommit_block():
for index_name, _ in reversed(_INDEXES):
op.drop_index(index_name, table_name=_TABLE_NAME, postgresql_concurrently=True)
return
for index_name, _ in reversed(_INDEXES):
op.drop_index(index_name, table_name=_TABLE_NAME)
+2
View File
@@ -101,6 +101,7 @@ from .model import (
UploadFile,
)
from .oauth import DatasourceOauthParamConfig, DatasourceProvider, OAuthAccessToken
from .onboarding import AccountStepByStepTourState
from .provider import (
LoadBalancingModelConfig,
Provider,
@@ -155,6 +156,7 @@ __all__ = [
"Account",
"AccountIntegrate",
"AccountStatus",
"AccountStepByStepTourState",
"AccountTrialAppRecord",
"Agent",
"AgentConfigDraft",
+59
View File
@@ -0,0 +1,59 @@
"""Account-level onboarding state models."""
from datetime import datetime
import sqlalchemy as sa
from sqlalchemy import DateTime, func
from sqlalchemy.orm import Mapped, mapped_column
from .base import TypeBase, gen_uuidv7_string
from .types import AdjustedJSON, StringUUID
class AccountStepByStepTourState(TypeBase):
"""Persistent account-level Step-by-step Tour state.
The tour is account-owned, with workspace IDs stored only as presentation
overrides. The first workspace is the workspace context where an eligible
account first asks for tour state; subsequent workspaces are opt-in only.
"""
__tablename__ = "account_step_by_step_tour_states"
__table_args__ = (
sa.PrimaryKeyConstraint("id", name="account_step_by_step_tour_state_pkey"),
sa.UniqueConstraint("account_id", name="account_step_by_step_tour_state_account_id_key"),
)
id: Mapped[str] = mapped_column(
StringUUID,
insert_default=gen_uuidv7_string,
default_factory=gen_uuidv7_string,
init=False,
)
account_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
first_workspace_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True, default=None)
skipped: Mapped[bool] = mapped_column(sa.Boolean, nullable=False, server_default=sa.text("false"), default=False)
completed_task_ids: Mapped[list[str]] = mapped_column(AdjustedJSON, nullable=False, default_factory=list)
manually_enabled_workspace_ids: Mapped[list[str]] = mapped_column(
AdjustedJSON,
nullable=False,
default_factory=list,
)
manually_disabled_workspace_ids: Mapped[list[str]] = mapped_column(
AdjustedJSON,
nullable=False,
default_factory=list,
)
created_at: Mapped[datetime] = mapped_column(
DateTime,
server_default=func.current_timestamp(),
nullable=False,
init=False,
)
updated_at: Mapped[datetime] = mapped_column(
DateTime,
server_default=func.current_timestamp(),
nullable=False,
init=False,
onupdate=func.current_timestamp(),
)
+2
View File
@@ -1476,6 +1476,8 @@ class WorkflowRunArchiveBundle(DefaultFieldsDCMixin, TypeBase):
name="workflow_run_archive_bundle_identity_uq",
),
sa.Index("workflow_run_archive_bundle_tenant_month_idx", "tenant_id", "year", "month"),
sa.Index("workflow_run_archive_bundle_month_id_idx", "year", "month", "id"),
sa.Index("workflow_run_archive_bundle_month_shard_id_idx", "year", "month", "shard", "id"),
)
tenant_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
+87
View File
@@ -7787,6 +7787,30 @@ Initiate OAuth login process
| ---- | ----------- | ------ |
| 200 | Success | **application/json**: [OAuthProviderTokenResponse](#oauthprovidertokenresponse)<br> |
### [GET] /onboarding/step-by-step-tour/state
Get account-level Step-by-step Tour state
#### Responses
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 200 | Success | **application/json**: [StepByStepTourStateResponse](#stepbysteptourstateresponse)<br> |
### [PATCH] /onboarding/step-by-step-tour/state
Update account-level Step-by-step Tour state
#### Request Body
| Required | Schema |
| -------- | ------ |
| Yes | **application/json**: [StepByStepTourStatePatchPayload](#stepbysteptourstatepatchpayload)<br> |
#### Responses
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 200 | Success | **application/json**: [StepByStepTourStateResponse](#stepbysteptourstateresponse)<br> |
### [DELETE] /rag/pipeline/customized/templates/{template_id}
#### Parameters
@@ -9639,6 +9663,27 @@ Bedrock retrieval test (internal use only)
| ---- | ----------- | ------ |
| 200 | Success | **application/json**: [TrialDatasetListResponse](#trialdatasetlistresponse)<br> |
### [POST] /trial-apps/{app_id}/files/upload
**Upload a file into the tenant that owns the trial app**
#### Parameters
| Name | Located in | Description | Required | Schema |
| ---- | ---------- | ----------- | -------- | ------ |
| app_id | path | | Yes | string (uuid) |
#### Request Body
| Required | Schema |
| -------- | ------ |
| Yes | **multipart/form-data**: { **"file"**: binary, **"source"**: string, <br>**Available values:** "datasets" }<br> |
#### Responses
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 201 | File uploaded successfully | **application/json**: [FileResponse](#fileresponse)<br> |
### [GET] /trial-apps/{app_id}/messages/{message_id}/suggested-questions
#### Parameters
@@ -9668,6 +9713,27 @@ Bedrock retrieval test (internal use only)
| ---- | ----------- | ------ |
| 200 | Success | **application/json**: [Parameters](#parameters)<br> |
### [POST] /trial-apps/{app_id}/remote-files/upload
**Upload a remote file into the tenant that owns the trial app**
#### Parameters
| Name | Located in | Description | Required | Schema |
| ---- | ---------- | ----------- | -------- | ------ |
| app_id | path | | Yes | string (uuid) |
#### Request Body
| Required | Schema |
| -------- | ------ |
| Yes | **application/json**: [RemoteFileUploadPayload](#remotefileuploadpayload)<br> |
#### Responses
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 201 | File uploaded successfully | **application/json**: [FileWithSignedUrl](#filewithsignedurl)<br> |
### [GET] /trial-apps/{app_id}/site
**Retrieve app site info**
@@ -21746,6 +21812,24 @@ Query parameters for listing snippet published workflows.
| paused | integer | | Yes |
| success | integer | | Yes |
#### StepByStepTourStatePatchPayload
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| action | string, <br>**Available values:** "complete_task", "disable_current_workspace", "enable_current_workspace", "skip", "uncomplete_task" | State update action<br>*Enum:* `"complete_task"`, `"disable_current_workspace"`, `"enable_current_workspace"`, `"skip"`, `"uncomplete_task"` | Yes |
| task_id | string | Task ID for task actions | No |
#### StepByStepTourStateResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| completed_task_ids | [ string, <br>**Available values:** "home", "integration", "knowledge", "studio" ] | | No |
| first_workspace_id | string | | No |
| manually_disabled_workspace_ids | [ string ] | | No |
| manually_enabled_workspace_ids | [ string ] | | No |
| skipped | boolean | | No |
| updated_at | string | | No |
#### Storage
| Name | Type | Description | Required |
@@ -21859,6 +21943,7 @@ The subscription constructor of the trigger provider
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| _is_collaborative | boolean | | No |
| conversation_variables | [ object ] | | No |
| environment_variables | [ object ] | | No |
| features | object | | Yes |
@@ -21898,6 +21983,7 @@ Model class for provider system configuration response.
| enable_learn_app | boolean, <br>**Default:** true | | Yes |
| enable_marketplace | boolean | | Yes |
| enable_social_oauth_login | boolean | | Yes |
| enable_step_by_step_tour | boolean | | Yes |
| enable_trial_app | boolean | | Yes |
| is_allow_create_workspace | boolean | | Yes |
| is_allow_register | boolean | | Yes |
@@ -22854,6 +22940,7 @@ in form definiton, or a variable while the workflow is running.
| ---- | ---- | ----------- | -------- |
| allow_email_code_login | boolean | | Yes |
| allow_email_password_login | boolean | | Yes |
| allow_public_access | boolean, <br>**Default:** true | | Yes |
| allow_sso | boolean | | Yes |
| enabled | boolean | | Yes |
| sso_config | [WebAppAuthSSOModel](#webappauthssomodel) | | Yes |
+2
View File
@@ -1577,6 +1577,7 @@ Default configuration for form inputs.
| enable_learn_app | boolean, <br>**Default:** true | | Yes |
| enable_marketplace | boolean | | Yes |
| enable_social_oauth_login | boolean | | Yes |
| enable_step_by_step_tour | boolean | | Yes |
| enable_trial_app | boolean | | Yes |
| is_allow_create_workspace | boolean | | Yes |
| is_allow_register | boolean | | Yes |
@@ -1642,6 +1643,7 @@ in form definiton, or a variable while the workflow is running.
| ---- | ---- | ----------- | -------- |
| allow_email_code_login | boolean | | Yes |
| allow_email_password_login | boolean | | Yes |
| allow_public_access | boolean, <br>**Default:** true | | Yes |
| allow_sso | boolean | | Yes |
| enabled | boolean | | Yes |
| sso_config | [WebAppAuthSSOModel](#webappauthssomodel) | | Yes |
+35
View File
@@ -0,0 +1,35 @@
"""Shared database fixtures for provider tests.
Provider tests live outside ``tests/unit_tests`` and therefore cannot use that
suite's SQLite fixtures. Keep this fixture scoped to ``providers`` so each
provider package can exercise real SQLAlchemy queries without a service
database or mocked sessions.
"""
from collections.abc import Iterator
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import Session
from models.base import TypeBase
@pytest.fixture
def sqlite3_session(request: pytest.FixtureRequest) -> Iterator[Session]:
"""Yield an isolated SQLite session with the parametrized model tables.
Pass the required ORM classes through indirect parametrization. The engine
is per-test so committed rows and identity-map state cannot leak between
provider packages.
"""
models: tuple[type[TypeBase], ...] = request.param
engine = create_engine("sqlite:///:memory:")
tables = [model.metadata.tables[model.__tablename__] for model in models]
TypeBase.metadata.create_all(engine, tables=tables)
try:
with Session(engine, expire_on_commit=False) as session:
yield session
finally:
engine.dispose()
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "dify-api"
version = "1.16.0-rc1"
version = "1.16.0"
requires-python = "~=3.12.0"
dependencies = [
@@ -3,7 +3,10 @@ from __future__ import annotations
import json
from typing import NotRequired, TypedDict, override
from redis.lock import Lock
from extensions.ext_redis import redis_client
from extensions.redis_names import serialize_redis_name
SESSION_STATE_TTL_SECONDS = 3600
SERVER_HEARTBEAT_TTL_SECONDS = 90
@@ -12,6 +15,31 @@ WORKFLOW_LEADER_PREFIX = "workflow_leader:"
WS_SID_MAP_PREFIX = "ws_sid_map:"
WS_SERVER_HEARTBEAT_PREFIX = "ws_server_heartbeat:"
WS_SERVER_SESSIONS_PREFIX = "ws_server_sessions:"
GRAPH_VIEW_STATE_LOCK_PREFIX = "workflow_graph_view_state_lock:"
GRAPH_VIEW_STATE_LOCK_TIMEOUT_SECONDS = 15
_UPDATE_SESSION_GRAPH_ACTIVE_LUA = """
local raw = redis.call('HGET', KEYS[1], ARGV[1])
if not raw then
return 0
end
local decoded_ok, session_info = pcall(cjson.decode, raw)
if not decoded_ok or type(session_info) ~= 'table' then
return 0
end
local incoming_sequence = tonumber(ARGV[3])
local current_sequence = tonumber(session_info.graph_active_sequence)
if current_sequence and incoming_sequence <= current_sequence then
return 0
end
session_info.graph_active = ARGV[2] == '1'
session_info.graph_active_sequence = incoming_sequence
redis.call('HSET', KEYS[1], ARGV[1], cjson.encode(session_info))
return 1
"""
class WorkflowSessionInfo(TypedDict):
@@ -156,6 +184,55 @@ class WorkflowCollaborationRepository:
return users
def get_session_info(self, workflow_id: str, sid: str) -> WorkflowSessionInfo | None:
raw = self._redis.hget(self.workflow_key(workflow_id), sid)
value = self._decode(raw)
if not value:
return None
try:
session_info = json.loads(value)
except (TypeError, json.JSONDecodeError):
return None
if not isinstance(session_info, dict):
return None
if "user_id" not in session_info or "username" not in session_info or "sid" not in session_info:
return None
user: WorkflowSessionInfo = {
"user_id": str(session_info["user_id"]),
"username": str(session_info["username"]),
"avatar": session_info.get("avatar"),
"sid": str(session_info["sid"]),
"connected_at": int(session_info.get("connected_at") or 0),
}
if isinstance(session_info.get("server_id"), str):
user["server_id"] = session_info["server_id"]
if isinstance(session_info.get("graph_active"), bool):
user["graph_active"] = session_info["graph_active"]
return user
def update_session_graph_active(self, workflow_id: str, sid: str, active: bool, sequence: int) -> bool:
"""Atomically apply a graph visibility update when its client sequence is newer."""
# RedisClientWrapper prefixes regular hash calls, but eval is delegated to the raw client.
workflow_key = serialize_redis_name(self.workflow_key(workflow_id))
result = self._redis.eval(
_UPDATE_SESSION_GRAPH_ACTIVE_LUA,
1,
workflow_key,
sid,
"1" if active else "0",
sequence,
)
return bool(result)
def graph_view_state_lock(self, workflow_id: str) -> Lock:
"""Serialize visibility state changes and their leader-election side effects."""
return self._redis.lock(
f"{GRAPH_VIEW_STATE_LOCK_PREFIX}{workflow_id}",
timeout=GRAPH_VIEW_STATE_LOCK_TIMEOUT_SECONDS,
)
def refresh_server_heartbeat(self, server_id: str) -> None:
self._redis.set(self.server_key(server_id), "1", ex=SERVER_HEARTBEAT_TTL_SECONDS)
+227 -18
View File
@@ -9,12 +9,13 @@ import sqlalchemy as sa
from sqlalchemy import and_, func, or_, select
from sqlalchemy.orm import aliased
from configs import dify_config
from core.app.entities.app_invoke_entities import InvokeFrom
from libs.helper import convert_datetime_to_date, escape_like_pattern, to_timestamp
from models.agent import WorkflowAgentNodeBinding
from models.enums import MessageStatus
from models.enums import CreatorUserRole, MessageStatus
from models.model import App, Conversation, Message
from models.workflow import WorkflowNodeExecutionModel, WorkflowRun
from models.workflow import WorkflowNodeExecutionModel, WorkflowRun, WorkflowType
@dataclass(frozen=True)
@@ -580,9 +581,26 @@ class AgentObservabilityService:
def _load_daily_statistics(
self, *, app: App, agent_id: str, params: AgentStatisticsQueryParams, source_filter: AgentSourceFilter
) -> list[dict[str, Any]]:
rows: list[dict[str, Any]] = []
if source_filter.kind in {"all", "webapp"}:
rows.extend(self._load_webapp_daily_statistics(app=app, params=params, source_filter=source_filter))
if source_filter.kind in {"all", "workflow"}:
rows.extend(
self._load_workflow_daily_statistics(
app=app,
agent_id=agent_id,
params=params,
source_filter=source_filter,
)
)
return self._merge_daily_statistics(rows)
def _load_webapp_daily_statistics(
self, *, app: App, params: AgentStatisticsQueryParams, source_filter: AgentSourceFilter
) -> list[dict[str, Any]]:
converted_created_at = convert_datetime_to_date("m.created_at")
message_scope = self._statistics_message_scope_sql(source_filter)
message_scope = self._statistics_webapp_message_scope_sql(source_filter)
sql_query = f"""SELECT
{converted_created_at} AS date,
COUNT(m.id) AS message_count,
@@ -602,12 +620,154 @@ WHERE
args: dict[str, Any] = {
"tz": params.timezone,
"app_id": app.id,
"tenant_id": app.tenant_id,
"agent_id": agent_id,
"debugger": InvokeFrom.DEBUGGER,
}
if source_filter.invoke_from is not None:
args["source"] = source_filter.invoke_from
if params.start:
sql_query += " AND m.created_at >= :start"
args["start"] = params.start
if params.end:
sql_query += " AND m.created_at < :end"
args["end"] = params.end
sql_query += " GROUP BY date ORDER BY date"
return [dict(row._mapping) for row in self._session.execute(sa.text(sql_query), args).all()]
@staticmethod
def _statistics_webapp_message_scope_sql(source_filter: AgentSourceFilter) -> str:
app_scope = "m.app_id = :app_id"
if source_filter.invoke_from is not None:
app_scope += " AND m.invoke_from = :source"
return app_scope
def _load_workflow_daily_statistics(
self,
*,
app: App,
agent_id: str,
params: AgentStatisticsQueryParams,
source_filter: AgentSourceFilter,
) -> list[dict[str, Any]]:
converted_run_created_at = convert_datetime_to_date("aru.created_at")
total_tokens = self._workflow_execution_metadata_numeric_sql(("total_tokens",), "BIGINT")
nested_total_tokens = self._workflow_execution_metadata_numeric_sql(
("agent_log", "agent_backend", "usage", "total_tokens"), "BIGINT"
)
total_price = self._workflow_execution_metadata_numeric_sql(("total_price",), "DECIMAL(65, 30)")
nested_total_price = self._workflow_execution_metadata_numeric_sql(
("agent_log", "agent_backend", "usage", "total_price"), "DECIMAL(65, 30)"
)
completion_tokens = self._workflow_execution_metadata_numeric_sql(
("agent_log", "agent_backend", "usage", "completion_tokens"), "BIGINT"
)
binding_filters = self._statistics_workflow_binding_filters_sql(source_filter)
run_date_filters = ""
args: dict[str, Any] = {
"tz": params.timezone,
"tenant_id": app.tenant_id,
"agent_id": agent_id,
"chat_workflow_type": WorkflowType.CHAT,
"end_user_role": CreatorUserRole.END_USER,
}
if source_filter.app_id:
args["source_app_id"] = source_filter.app_id
if source_filter.workflow_id:
args["workflow_id"] = source_filter.workflow_id
if source_filter.workflow_version:
args["workflow_version"] = source_filter.workflow_version
if source_filter.node_id:
args["node_id"] = source_filter.node_id
if params.start:
run_date_filters += " AND wr.created_at >= :start"
args["start"] = params.start
if params.end:
run_date_filters += " AND wr.created_at < :end"
args["end"] = params.end
run_query = f"""WITH agent_run_usage AS (
SELECT
wr.id,
wr.created_by_role,
wr.created_by,
wr.created_at,
COALESCE(SUM(COALESCE({total_tokens}, {nested_total_tokens}, 0)), 0) AS token_count,
COALESCE(SUM(COALESCE({total_price}, {nested_total_price}, 0)), 0) AS total_price,
COALESCE(SUM(COALESCE(wne.elapsed_time, 0)), 0) AS latency,
COALESCE(SUM(COALESCE({completion_tokens}, 0)), 0) AS answer_tokens
FROM workflow_runs wr
JOIN workflow_agent_node_bindings wanb
ON wanb.tenant_id = :tenant_id
AND wanb.agent_id = :agent_id
AND wanb.app_id = wr.app_id
AND wanb.workflow_id = wr.workflow_id
AND wanb.workflow_version = wr.version
{binding_filters}
JOIN workflow_node_executions wne
ON wne.workflow_run_id = wr.id
AND wne.node_id = wanb.node_id
WHERE wr.type != :chat_workflow_type{run_date_filters}
GROUP BY wr.id, wr.created_by_role, wr.created_by, wr.created_at
)
SELECT
{converted_run_created_at} AS date,
COUNT(aru.id) AS message_count,
COUNT(aru.id) AS conversation_count,
COUNT(DISTINCT CASE
WHEN aru.created_by_role = :end_user_role THEN aru.created_by
ELSE NULL
END) AS end_user_count,
COALESCE(SUM(aru.token_count), 0) AS token_count,
COALESCE(SUM(aru.total_price), 0) AS total_price,
COALESCE(AVG(aru.latency), 0) AS avg_latency,
COALESCE(SUM(aru.latency), 0) AS latency_sum,
COALESCE(SUM(aru.answer_tokens), 0) AS answer_tokens,
0 AS like_count
FROM agent_run_usage aru
GROUP BY date
ORDER BY date"""
rows = [dict(row._mapping) for row in self._session.execute(sa.text(run_query), args).all()]
rows.extend(
self._load_workflow_chat_daily_context(
app=app,
agent_id=agent_id,
params=params,
source_filter=source_filter,
)
)
return self._merge_daily_statistics(rows)
def _load_workflow_chat_daily_context(
self,
*,
app: App,
agent_id: str,
params: AgentStatisticsQueryParams,
source_filter: AgentSourceFilter,
) -> list[dict[str, Any]]:
converted_created_at = convert_datetime_to_date("m.created_at")
workflow_scope = self._statistics_workflow_message_scope_sql(source_filter)
sql_query = f"""SELECT
{converted_created_at} AS date,
COUNT(m.id) AS message_count,
COUNT(DISTINCT m.conversation_id) AS conversation_count,
COUNT(DISTINCT m.from_end_user_id) AS end_user_count,
COALESCE(SUM(COALESCE(m.message_tokens, 0) + COALESCE(m.answer_tokens, 0)), 0) AS token_count,
COALESCE(SUM(COALESCE(m.total_price, 0)), 0) AS total_price,
COALESCE(AVG(m.provider_response_latency), 0) AS avg_latency,
COALESCE(SUM(m.provider_response_latency), 0) AS latency_sum,
COALESCE(SUM(m.answer_tokens), 0) AS answer_tokens,
COUNT(mf.id) AS like_count
FROM messages m
LEFT JOIN message_feedbacks mf
ON mf.message_id = m.id AND mf.rating = 'like'
WHERE
{workflow_scope}"""
args: dict[str, Any] = {
"tz": params.timezone,
"tenant_id": app.tenant_id,
"agent_id": agent_id,
"chat_workflow_type": WorkflowType.CHAT,
}
if source_filter.app_id:
args["source_app_id"] = source_filter.app_id
if source_filter.workflow_id:
@@ -627,10 +787,7 @@ WHERE
return [dict(row._mapping) for row in self._session.execute(sa.text(sql_query), args).all()]
@staticmethod
def _statistics_message_scope_sql(source_filter: AgentSourceFilter) -> str:
app_scope = "m.app_id = :app_id"
if source_filter.invoke_from is not None:
app_scope += " AND m.invoke_from = :source"
def _statistics_workflow_binding_filters_sql(source_filter: AgentSourceFilter) -> str:
workflow_binding_filters = []
if source_filter.app_id:
workflow_binding_filters.append("wanb.app_id = :source_app_id")
@@ -640,8 +797,12 @@ WHERE
workflow_binding_filters.append("wanb.workflow_version = :workflow_version")
if source_filter.node_id:
workflow_binding_filters.append("wanb.node_id = :node_id")
extra_workflow_filters = f"AND {' AND '.join(workflow_binding_filters)}" if workflow_binding_filters else ""
workflow_scope = f"""m.workflow_run_id IS NOT NULL
return f"AND {' AND '.join(workflow_binding_filters)}" if workflow_binding_filters else ""
@classmethod
def _statistics_workflow_message_scope_sql(cls, source_filter: AgentSourceFilter) -> str:
binding_filters = cls._statistics_workflow_binding_filters_sql(source_filter)
return f"""m.workflow_run_id IS NOT NULL
AND EXISTS (
SELECT 1
FROM workflow_runs wr
@@ -651,17 +812,65 @@ WHERE
AND wanb.app_id = wr.app_id
AND wanb.workflow_id = wr.workflow_id
AND wanb.workflow_version = wr.version
{extra_workflow_filters}
{binding_filters}
JOIN workflow_node_executions wne
ON wne.workflow_run_id = wr.id
AND wne.node_id = wanb.node_id
WHERE wr.id = m.workflow_run_id
AND wr.type = :chat_workflow_type
)"""
if source_filter.kind == "webapp":
return app_scope
if source_filter.kind == "workflow":
return workflow_scope
return f"(({app_scope}) OR ({workflow_scope}))"
@staticmethod
def _workflow_execution_metadata_numeric_sql(path: tuple[str, ...], numeric_type: str) -> str:
if dify_config.DB_TYPE == "postgresql":
json_path = ",".join(path)
value = f"CAST(wne.execution_metadata AS JSONB) #>> '{{{json_path}}}'"
return f"CAST(NULLIF({value}, '') AS {numeric_type})"
if dify_config.DB_TYPE in {"mysql", "oceanbase", "seekdb"}:
json_path = "$." + ".".join(path)
mysql_numeric_type = "UNSIGNED" if numeric_type == "BIGINT" else numeric_type
value = f"JSON_UNQUOTE(JSON_EXTRACT(wne.execution_metadata, '{json_path}'))"
return f"CAST(NULLIF(NULLIF({value}, ''), 'null') AS {mysql_numeric_type})"
raise NotImplementedError(f"Unsupported database type: {dify_config.DB_TYPE}")
@staticmethod
def _merge_daily_statistics(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
merged: dict[Any, dict[str, Any]] = {}
weighted_latency: dict[Any, float] = {}
for row in rows:
date = row["date"]
target = merged.setdefault(
date,
{
"date": date,
"message_count": 0,
"conversation_count": 0,
"end_user_count": 0,
"token_count": 0,
"total_price": Decimal(0),
"avg_latency": 0.0,
"latency_sum": 0.0,
"answer_tokens": 0,
"like_count": 0,
},
)
message_count = int(row.get("message_count") or 0)
target["message_count"] += message_count
target["conversation_count"] += int(row.get("conversation_count") or 0)
target["end_user_count"] += int(row.get("end_user_count") or 0)
target["token_count"] += int(row.get("token_count") or 0)
target["total_price"] += Decimal(str(row.get("total_price") or 0))
target["latency_sum"] += float(row.get("latency_sum") or 0)
target["answer_tokens"] += int(row.get("answer_tokens") or 0)
target["like_count"] += int(row.get("like_count") or 0)
weighted_latency[date] = (
weighted_latency.get(date, 0.0) + float(row.get("avg_latency") or 0) * message_count
)
for date, row in merged.items():
message_count = int(row["message_count"])
row["avg_latency"] = weighted_latency[date] / message_count if message_count else 0.0
return sorted(merged.values(), key=lambda row: str(row["date"]))
@staticmethod
def _build_charts(rows: list[dict[str, Any]]) -> dict[str, list[dict[str, Any]]]:
+24 -12
View File
@@ -408,6 +408,7 @@ class AppService:
default_model_config = app_template.get("model_config")
default_model_config = default_model_config.copy() if default_model_config else None
if default_model_config and "model" in default_model_config:
default_model_dict = default_model_config["model"]
# get model provider
model_manager = ModelManager.for_tenant(tenant_id=account.current_tenant_id or "")
@@ -422,7 +423,7 @@ class AppService:
logger.exception("Get default model instance failed, tenant_id: %s", tenant_id)
model_instance = None
if model_instance:
if model_instance is not None:
if (
model_instance.model_name == default_model_config["model"]["name"]
and model_instance.provider == default_model_config["model"]["provider"]
@@ -430,17 +431,28 @@ class AppService:
default_model_dict = default_model_config["model"]
else:
llm_model = cast(LargeLanguageModel, model_instance.model_type_instance)
model_schema = llm_model.get_model_schema(model_instance.model_name, model_instance.credentials)
if model_schema is None:
raise ValueError(f"model schema not found for model {model_instance.model_name}")
default_model_dict = {
"provider": model_instance.provider,
"name": model_instance.model_name,
"mode": model_schema.model_properties.get(ModelPropertyKey.MODE),
"completion_params": {},
}
else:
try:
model_schema = llm_model.get_model_schema(model_instance.model_name, model_instance.credentials)
if model_schema is None:
raise ValueError(f"model schema not found for model {model_instance.model_name}")
except Exception:
# A removed provider model must not prevent creating an app.
logger.warning(
"Default model schema is unavailable, tenant_id: %s, provider: %s, model: %s",
tenant_id,
model_instance.provider,
model_instance.model_name,
exc_info=True,
)
model_instance = None
else:
default_model_dict = {
"provider": model_instance.provider,
"name": model_instance.model_name,
"mode": model_schema.model_properties.get(ModelPropertyKey.MODE),
"completion_params": {},
}
if model_instance is None:
try:
provider, model = model_manager.get_default_provider_model_name(
tenant_id=account.current_tenant_id or "", model_type=ModelType.LLM
+15 -4
View File
@@ -739,8 +739,13 @@ class DatasetService:
dataset.id, external_knowledge_id, external_knowledge_api_id, session
)
# Commit changes to database
session.commit()
# Flush changes to the database without closing the caller-managed
# transaction. This helper receives a session opened by the caller
# (`with Session(...) as session`); calling commit() here closed that
# context manager early and raised
# sqlalchemy.exc.InvalidRequestError: Can't operate on closed transaction
# (#39191).
session.flush()
return dataset
@@ -810,9 +815,15 @@ class DatasetService:
if data.get("icon_info"):
filtered_data["icon_info"] = data.get("icon_info")
# Update dataset in database
# Update dataset in database. Use flush() rather than commit() so the
# caller-managed transaction (opened with `with Session(...) as session`)
# stays open for subsequent operations — _update_pipeline_knowledge_base
# node data and any caller follow-ups run on the same session. Calling
# commit() here closed the context manager early and raised
# sqlalchemy.exc.InvalidRequestError: Can't operate on closed transaction
# (#39191).
session.execute(update(Dataset).where(Dataset.id == dataset.id).values(**filtered_data))
session.commit()
session.flush()
# Reload dataset to get updated values
session.refresh(dataset)
-13
View File
@@ -317,9 +317,6 @@ _LEGACY_WORKSPACE_OWNER_KEYS: list[str] = [
"credential.use",
"credential.create",
"credential.manage",
"billing.view",
"billing.subscription.manage",
"billing.manage",
"app.acl.preview",
"app_library.access",
"app.create_and_management",
@@ -349,9 +346,6 @@ _LEGACY_WORKSPACE_ADMIN_KEYS: list[str] = [
"credential.use",
"credential.create",
"credential.manage",
"billing.view",
"billing.subscription.manage",
"billing.manage",
"app_library.access",
"app.create_and_management",
"app.tag.manage",
@@ -377,9 +371,6 @@ _LEGACY_WORKSPACE_EDITOR_KEYS: list[str] = [
"dataset.external.connect",
"snippets.create_and_modify",
"tool.manage",
"billing.view",
"billing.subscription.manage",
"billing.manage",
]
_LEGACY_WORKSPACE_NORMAL_KEYS: list[str] = [
@@ -387,9 +378,6 @@ _LEGACY_WORKSPACE_NORMAL_KEYS: list[str] = [
"plugin.install",
"credential.use",
"app_library.access",
"billing.view",
"billing.subscription.manage",
"billing.manage",
]
_LEGACY_WORKSPACE_DATASET_OPERATOR_KEYS: list[str] = [
@@ -805,7 +793,6 @@ class RBACService:
data = _inner_call(
"GET",
f"{_INNER_PREFIX}/role-permissions/catalog",
params={"billing_enabled": dify_config.BILLING_ENABLED},
tenant_id=tenant_id,
account_id=account_id,
)
+4
View File
@@ -100,6 +100,7 @@ class WebAppAuthModel(FeatureResponseModel):
sso_config: WebAppAuthSSOModel = WebAppAuthSSOModel()
allow_email_code_login: bool = False
allow_email_password_login: bool = False
allow_public_access: bool = True
class KnowledgePipeline(FeatureResponseModel):
@@ -183,6 +184,7 @@ class SystemFeatureModel(FeatureResponseModel):
enable_trial_app: bool = False
enable_explore_banner: bool = False
enable_learn_app: bool = True
enable_step_by_step_tour: bool = False
rbac_enabled: bool = False
@@ -285,6 +287,8 @@ class FeatureService:
system_features.enable_trial_app = dify_config.ENABLE_TRIAL_APP
system_features.enable_explore_banner = dify_config.ENABLE_EXPLORE_BANNER
system_features.enable_learn_app = dify_config.ENABLE_LEARN_APP
system_features.webapp_auth.allow_public_access = dify_config.WEBAPP_PUBLIC_ACCESS_ENABLED
system_features.enable_step_by_step_tour = dify_config.ENABLE_STEP_BY_STEP_TOUR
@classmethod
def _fulfill_trial_models_from_env(cls) -> list[str]:
+358
View File
@@ -0,0 +1,358 @@
"""Transport-only forwarding for the explicitly enabled KnowledgeFS Console operations.
KnowledgeFS owns the request and response contract. This module binds short-lived
account and workspace identities, enforces Dify's coarse workspace policy, and
normalizes transport failures. Dify deliberately maintains a small product-facing
operation registry instead of exposing the full upstream OpenAPI surface. The
dedicated request path uses Dify's shared SSRF policy, never follows redirects,
bounds buffered responses, and rejects compressed responses.
"""
from __future__ import annotations
from collections.abc import Iterable, Mapping
from datetime import UTC, datetime, timedelta
from http import HTTPStatus
from typing import Final, Literal, NamedTuple, Protocol
import httpx
import jwt
from configs import dify_config
from core.helper import ssrf_proxy
from core.rbac import RBACPermission, RBACResourceScope
from core.tools.errors import ToolSSRFError
from models import Account
from services.enterprise.rbac_service import RBACService
type KnowledgeFSMethod = Literal["DELETE", "GET", "PATCH", "POST", "PUT"]
type KnowledgeFSResponseKind = Literal["binary", "buffered", "stream"]
type KnowledgeFSRequiredScope = Literal["knowledge-spaces:read", "knowledge-spaces:write"]
_JWT_AUDIENCE = "knowledge-fs"
_JWT_ISSUER = "dify"
_JWT_TTL_SECONDS = 60
_MAX_BUFFERED_RESPONSE_BYTES = 1024 * 1024
class KnowledgeFSOperation(NamedTuple):
operation_id: str
method: KnowledgeFSMethod
path: str
response_kind: KnowledgeFSResponseKind
required_scope: KnowledgeFSRequiredScope
rbac_permission: RBACPermission
requires_dataset_editor: bool
max_response_bytes: int
request_headers: tuple[str, ...]
response_headers: tuple[str, ...]
response_media_types: tuple[str, ...]
KNOWLEDGE_FS_CONSOLE_OPERATIONS: Final[tuple[KnowledgeFSOperation, ...]] = (
KnowledgeFSOperation(
operation_id="listKnowledgeSpaces",
method="GET",
path="knowledge-spaces",
response_kind="buffered",
required_scope="knowledge-spaces:read",
rbac_permission=RBACPermission.DATASET_READONLY,
requires_dataset_editor=False,
max_response_bytes=1_048_576,
request_headers=("x-trace-id",),
response_headers=("x-trace-id",),
response_media_types=("application/json",),
),
KnowledgeFSOperation(
operation_id="createKnowledgeSpace",
method="POST",
path="knowledge-spaces",
response_kind="buffered",
required_scope="knowledge-spaces:write",
rbac_permission=RBACPermission.DATASET_CREATE_AND_MANAGEMENT,
requires_dataset_editor=True,
max_response_bytes=1_048_576,
request_headers=("x-trace-id",),
response_headers=("x-trace-id",),
response_media_types=("application/json",),
),
)
class KnowledgeFSUpstreamResponse(NamedTuple):
response: httpx.Response
response_kind: KnowledgeFSResponseKind
operation: KnowledgeFSOperation
class _RequestHeaders(Protocol):
def items(self) -> Iterable[tuple[str, str]]: ...
class KnowledgeFSConfigurationError(RuntimeError):
"""KnowledgeFS is incompletely configured or blocked by outbound policy."""
class KnowledgeFSTimeoutError(RuntimeError):
"""KnowledgeFS exceeded the configured request timeout."""
class KnowledgeFSTransportError(RuntimeError):
"""KnowledgeFS could not be reached or returned a response outside safety bounds."""
class KnowledgeFSRouteNotAllowedError(RuntimeError):
"""The requested path is outside the Console-visible KnowledgeFS surface."""
class KnowledgeFSAccessDeniedError(RuntimeError):
"""The Dify account lacks the workspace permission required by the operation."""
def authorize_knowledge_fs_request(
*,
account: Account,
tenant_id: str,
operation: KnowledgeFSOperation,
) -> None:
"""Enforce Dify's workspace policy before KFS performs resource authorization.
Args:
account: Authenticated Dify account with its current workspace role.
tenant_id: Current Dify workspace identifier.
operation: Dify-maintained KnowledgeFS operation and policy metadata.
Raises:
KnowledgeFSAccessDeniedError: The account lacks a required legacy or enterprise permission.
"""
if operation.requires_dataset_editor and not account.is_dataset_editor:
raise KnowledgeFSAccessDeniedError("KnowledgeFS mutations require dataset edit access")
if not RBACService.CheckAccess.check(
tenant_id,
account.id,
scene=operation.rbac_permission.value,
resource_type=RBACResourceScope.DATASET.value,
):
raise KnowledgeFSAccessDeniedError("KnowledgeFS operation is denied by workspace RBAC")
def proxy_knowledge_fs_request(
*,
account: Account,
method: KnowledgeFSMethod,
path: str,
tenant_id: str,
accept: str | None = None,
content_type: str | None = None,
query: bytes | None = None,
body: bytes | None = None,
request_headers: _RequestHeaders | None = None,
) -> KnowledgeFSUpstreamResponse:
"""Authorize and forward one allowlisted KnowledgeFS request as a single use case."""
operation = get_knowledge_fs_operation(method, path)
authorize_knowledge_fs_request(
account=account,
tenant_id=tenant_id,
operation=operation,
)
incoming_request_headers = {name.lower(): value for name, value in (request_headers or {}).items()}
contract_request_headers = {
name: incoming_request_headers[name] for name in operation.request_headers if name in incoming_request_headers
}
return _forward_knowledge_fs_request(
account_id=account.id,
method=method,
path=path,
tenant_id=tenant_id,
accept=accept,
content_type=content_type,
query=query,
body=body,
request_headers=contract_request_headers,
)
def _forward_knowledge_fs_request(
*,
account_id: str,
method: KnowledgeFSMethod,
path: str,
tenant_id: str,
accept: str | None = None,
content_type: str | None = None,
query: bytes | None = None,
body: bytes | None = None,
request_headers: Mapping[str, str] | None = None,
) -> KnowledgeFSUpstreamResponse:
"""Forward one fixed-route request without parsing its KnowledgeFS payload.
Args:
account_id: Current Dify account used as the KFS member identity.
method: Allowlisted upstream HTTP method.
path: Relative KnowledgeFS path under an allowlisted product surface.
tenant_id: Current Dify workspace used as the KFS tenant identity.
accept: Original Accept header, when present.
content_type: Original request Content-Type header, when present.
query: Original encoded query string from the Console request.
body: Original request body, when present.
request_headers: Contract-declared request headers forwarded by the Console adapter.
Returns:
The KnowledgeFS response and its actual transport kind. Non-success responses are buffered.
Raises:
KnowledgeFSConfigurationError: The connection is incomplete or blocked by outbound policy.
KnowledgeFSRouteNotAllowedError: The path is outside the allowlisted product surface.
KnowledgeFSTimeoutError: KnowledgeFS exceeds the configured timeout.
KnowledgeFSTransportError: The request fails or its response cannot be safely bounded.
Each request is bound to stable Dify account and workspace principals with a short expiration.
"""
operation = get_knowledge_fs_operation(method, path)
base_url = dify_config.KNOWLEDGE_FS_BASE_URL
jwt_secret = dify_config.KNOWLEDGE_FS_JWT_SECRET
if base_url is None or jwt_secret is None:
raise KnowledgeFSConfigurationError("KnowledgeFS connection configuration is incomplete")
now = datetime.now(UTC)
token = jwt.encode(
{
"aud": _JWT_AUDIENCE,
"caller_kind": "interactive",
"dify_account_id": f"dify-account:{account_id}",
"exp": now + timedelta(seconds=_JWT_TTL_SECONDS),
"iat": now,
"iss": _JWT_ISSUER,
"scopes": [operation.required_scope],
"sub": f"dify-workspace:{tenant_id}",
"tenant_id": tenant_id,
},
jwt_secret.get_secret_value(),
algorithm="HS256",
)
headers = {
"Accept": accept or "application/json",
"Accept-Encoding": "identity",
"Authorization": f"Bearer {token}",
}
if body is not None:
headers["Content-Type"] = content_type or "application/json"
allowed_request_headers = set(operation.request_headers)
for name, value in (request_headers or {}).items():
normalized_name = name.lower()
if normalized_name not in allowed_request_headers:
raise KnowledgeFSRouteNotAllowedError("KnowledgeFS request header is not allowed")
headers[normalized_name] = value
try:
upstream_url = httpx.URL(f"{base_url}/").join(operation.path)
response = ssrf_proxy.make_request(
method=operation.method,
url=str(upstream_url),
params=query,
content=body,
headers=headers,
timeout=dify_config.KNOWLEDGE_FS_TIMEOUT_SECONDS,
follow_redirects=False,
max_retries=0,
stream_response=True,
)
response_kind = _classify_response(operation, response)
if response_kind == "stream":
content_encoding = response.headers.get("content-encoding", "identity").strip().lower()
if content_encoding not in {"", "identity"}:
response.close()
raise KnowledgeFSTransportError("KnowledgeFS streaming response used an unsupported encoding")
_set_response_read_timeout(response, dify_config.KNOWLEDGE_FS_SSE_READ_TIMEOUT_SECONDS)
return KnowledgeFSUpstreamResponse(response, response_kind, operation)
max_response_bytes = (
operation.max_response_bytes
if HTTPStatus.OK <= response.status_code < HTTPStatus.MULTIPLE_CHOICES
else _MAX_BUFFERED_RESPONSE_BYTES
)
buffered_response = ssrf_proxy.buffer_response(response, max_response_bytes=max_response_bytes)
if buffered_response.content and not buffered_response.headers.get("content-type", "").strip():
buffered_response.close()
raise KnowledgeFSTransportError("KnowledgeFS buffered response used an unsupported media type")
return KnowledgeFSUpstreamResponse(buffered_response, response_kind, operation)
except ssrf_proxy.ResponseLimitError as exc:
raise KnowledgeFSTransportError("KnowledgeFS response violated the proxy limit") from exc
except ToolSSRFError as exc:
raise KnowledgeFSConfigurationError("KnowledgeFS origin was blocked by outbound policy") from exc
except httpx.TimeoutException as exc:
raise KnowledgeFSTimeoutError("KnowledgeFS request timed out") from exc
except httpx.RequestError as exc:
raise KnowledgeFSTransportError("KnowledgeFS transport request failed") from exc
def get_knowledge_fs_operation(method: KnowledgeFSMethod, path: str) -> KnowledgeFSOperation:
"""Resolve an exact operation and its transport/access contract metadata."""
for operation in KNOWLEDGE_FS_CONSOLE_OPERATIONS:
if method == operation.method and _matches_route_template(operation.path, path):
return operation._replace(path=path)
raise KnowledgeFSRouteNotAllowedError("KnowledgeFS route is not allowed")
def _classify_response(operation: KnowledgeFSOperation, response: httpx.Response) -> KnowledgeFSResponseKind:
"""Resolve the actual response kind from status and Content-Type before reading its body."""
content_type = response.headers.get("content-type", "").partition(";")[0].strip().lower()
is_success = HTTPStatus.OK <= response.status_code < HTTPStatus.MULTIPLE_CHOICES
if not is_success:
if content_type and not _is_json_content_type(content_type):
response.close()
raise KnowledgeFSTransportError("KnowledgeFS error response used an unsupported media type")
return "buffered"
if operation.response_kind == "stream":
if content_type != "text/event-stream":
response.close()
raise KnowledgeFSTransportError("KnowledgeFS stream response used an unsupported media type")
return "stream"
if operation.response_kind == "binary":
if content_type not in operation.response_media_types:
response.close()
raise KnowledgeFSTransportError("KnowledgeFS binary response used an unsupported media type")
return "binary"
if content_type and not _is_json_content_type(content_type):
response.close()
raise KnowledgeFSTransportError("KnowledgeFS buffered response used an unsupported media type")
return "buffered"
def _is_json_content_type(content_type: str) -> bool:
return content_type == "application/json" or content_type.endswith("+json")
def _set_response_read_timeout(response: httpx.Response, timeout_seconds: float | None) -> None:
"""Set the body-read timeout after headers identify a valid SSE response."""
try:
request = response.request
except RuntimeError:
return
timeout = request.extensions.get("timeout")
if isinstance(timeout, dict):
timeout["read"] = timeout_seconds
def _matches_route_template(template: str, path: str) -> bool:
"""Match path parameters without permitting encoded or traversal-like segments."""
template_segments = template.split("/")
path_segments = path.split("/")
if len(template_segments) != len(path_segments):
return False
for template_segment, path_segment in zip(template_segments, path_segments, strict=True):
if template_segment.startswith("{") and template_segment.endswith("}"):
if (
not path_segment
or path_segment in {".", ".."}
or "\\" in path_segment
or "%" in path_segment
or "?" in path_segment
or "#" in path_segment
):
return False
continue
if template_segment != path_segment:
return False
return True
@@ -601,6 +601,7 @@ class RagPipelineService:
with sessionmaker(bind=db.engine).begin() as session:
draft_var_saver = DraftVariableSaver(
session=session,
tenant_id=pipeline.tenant_id,
app_id=pipeline.id,
node_id=workflow_node_execution.node_id,
node_type=workflow_node_execution.node_type,
@@ -1391,6 +1392,7 @@ class RagPipelineService:
with sessionmaker(bind=db.engine).begin() as session:
draft_var_saver = DraftVariableSaver(
session=session,
tenant_id=pipeline.tenant_id,
app_id=pipeline.id,
node_id=workflow_node_execution_db_model.node_id,
node_type=workflow_node_execution_db_model.node_type,
@@ -6,7 +6,10 @@ This service archives workflow run logs for paid plan users older than the confi
Archive V2 writes bundle-level Parquet objects. A bundle contains many workflow runs and their related table rows.
Bundle metadata lives in the object-store manifest as the recoverable source of truth. Completed bundles are also
mirrored into a small database index so console listing and download jobs do not list object storage online.
published into a small database catalog so console listing, download, and maintenance jobs do not list object storage
online. Archive success requires that catalog publication to commit. A retry reconciles only its known manifest key
against an already-published shard index; a missing shard index fails closed rather than rebuilding it by scanning
historical manifests.
Archive campaigns should use fixed absolute UTC windows for every tenant-prefix/shard execution. Relative windows are
evaluated at process start and are not safe for multi-day rollout because each command would scan a different window.
@@ -639,7 +642,6 @@ class WorkflowRunArchiver:
if storage is None:
raise ArchiveStorageNotConfiguredError("Archive storage not configured")
if storage.object_exists(self._get_manifest_object_key(identity)):
self._write_bundle_index(storage, identity)
self._sync_existing_bundle_index(session, storage, identity)
result.success = True
result.skipped = True
@@ -651,6 +653,8 @@ class WorkflowRunArchiver:
runs = [run for run in runs if run.id not in archived_run_ids]
result.skipped_run_count = original_run_count - len(runs)
if not runs:
# Historical catalog rows are a rollout precondition. New bundles commit their catalog row before
# publishing the shard index, so this idempotency path never needs to rescan prior manifests.
result.run_count = 0
result.success = True
result.skipped = True
@@ -663,7 +667,6 @@ class WorkflowRunArchiver:
result.object_prefix = identity.object_prefix
result.run_count = len(runs)
if storage.object_exists(self._get_manifest_object_key(identity)):
self._write_bundle_index(storage, identity)
self._sync_existing_bundle_index(session, storage, identity)
result.success = True
result.skipped = True
@@ -695,10 +698,10 @@ class WorkflowRunArchiver:
for table_name, payload in table_payloads.items():
storage.put_object(self._get_table_object_key(identity, table_name), payload)
storage.put_object(self._get_manifest_object_key(identity), manifest_data)
self._merge_bundle_manifest_into_index(storage, identity, [run.id for run in runs])
manifest = decode_archive_bundle_manifest(manifest_data)
upsert_archive_bundle_index_from_manifest(session, manifest, len(manifest_data))
session.commit()
self._merge_bundle_manifest_into_index(storage, identity, [run.id for run in runs])
logger.info(
"Archived workflow run bundle %s: tenant=%s runs=%s tables=%s object_prefix=%s",
@@ -730,16 +733,18 @@ class WorkflowRunArchiver:
storage: ArchiveStorage,
identity: ArchiveBundleIdentity,
) -> None:
"""Best-effort DB index sync for a bundle whose manifest already exists in archive storage."""
"""Publish a known manifest to the DB catalog and reconcile its existing shard index."""
manifest_key = self._get_manifest_object_key(identity)
try:
manifest_data = storage.get_object(manifest_key)
manifest = decode_archive_bundle_manifest(manifest_data)
upsert_archive_bundle_index_from_manifest(session, manifest, len(manifest_data))
session.commit()
except Exception:
session.rollback()
logger.warning("Failed to sync workflow archive bundle index for %s", manifest_key, exc_info=True)
manifest_data = storage.get_object(manifest_key)
manifest = decode_archive_bundle_manifest(manifest_data)
upsert_archive_bundle_index_from_manifest(session, manifest, len(manifest_data))
session.commit()
self._merge_bundle_manifest_into_index(
storage,
identity,
manifest["run_ids"],
require_existing_index=True,
)
def _lock_runs_for_archive(
self,
@@ -1029,10 +1034,20 @@ class WorkflowRunArchiver:
storage: ArchiveStorage,
identity: ArchiveBundleIdentity,
run_ids: Sequence[str],
*,
require_existing_index: bool = False,
) -> ArchiveBundleIndexDict:
"""
Merge one bundle into its shard index.
Retries for a known manifest set ``require_existing_index`` so they never create a partial index by scanning
or overwriting a shard whose historical entries cannot be proven from this one manifest.
"""
index_key = self._get_index_object_key(identity)
if storage.object_exists(index_key):
index = self._load_bundle_index(storage, identity)
elif require_existing_index:
raise RuntimeError(f"archive shard index missing while reconciling known manifest: {index_key}")
else:
index = self._build_bundle_index(storage, identity)
@@ -1,14 +1,16 @@
"""
Maintain V2 workflow-run archive bundles.
Archive V2 keeps object-store manifests as the recoverable bundle source of truth. This maintenance module still
discovers delete/restore targets by listing `manifest.json` objects and uses object-store marker files for
delete/restore state. The separate database bundle index is intended for console listing and download jobs, not as the
source of truth for destructive maintenance.
Archive V2 keeps object-store manifests as the recoverable bundle source of truth. Delete and restore discover a
bounded page of candidates from `workflow_run_archive_bundles`, optionally restricted to one exact archive shard, then
construct each immutable manifest key from the catalog identity. They never list the object-store namespace.
Object-store marker files keep delete/restore idempotent, while the caller persists a non-dry-run cursor only after a
candidate succeeds.
Each bundle is processed in its own database transaction. A failed bundle leaves source rows unchanged unless the
transaction has already committed; marker handling makes the next run able to reconcile the common committed-but-marker
not-updated case.
not-updated case. Restore never skips a bundle with a missing deleted marker when deletion started or source rows have
drifted, so an external cursor cannot pass an interrupted delete.
"""
import datetime
@@ -38,6 +40,7 @@ from models.workflow import (
WorkflowPause,
WorkflowPauseReason,
WorkflowRun,
WorkflowRunArchiveBundle,
)
from services.retention.workflow_run.constants import (
ARCHIVE_BUNDLE_DELETE_STARTED_MARKER_NAME,
@@ -51,7 +54,6 @@ from services.retention.workflow_run.constants import (
logger = logging.getLogger(__name__)
_ARCHIVE_ROOT_PREFIX = "workflow-runs/v2/"
_CHUNK_SIZE = 5_000
@@ -84,11 +86,28 @@ class BundleManifest(TypedDict):
@dataclass(frozen=True)
class BundleReference:
"""Object-store reference for one V2 archive bundle."""
class ArchiveBundleCatalogEntry:
"""Immutable catalog identity and manifest-derived metrics for one V2 archive bundle."""
catalog_id: str
tenant_id: str
year: int
month: int
shard: str
bundle_id: str
workflow_run_count: int
row_count: int
archive_bytes: int
@dataclass(frozen=True)
class BundleReference:
"""Verified object-store reference for one catalog candidate."""
catalog: ArchiveBundleCatalogEntry
object_prefix: str
manifest_key: str
manifest_size_bytes: int
manifest: BundleManifest
@@ -96,6 +115,7 @@ class BundleReference:
class BundleOperationResult:
"""Result for one V2 bundle delete or restore operation."""
catalog_id: str
bundle_id: str
tenant_id: str
object_prefix: str
@@ -128,6 +148,8 @@ class BundleOperationSummary:
archive_bytes: int = 0
elapsed_time: float = 0.0
validation_time: float = 0.0
next_catalog_id: str | None = None
preview_next_catalog_id: str | None = None
table_counts: dict[str, int] = field(default_factory=dict)
results: list[BundleOperationResult] = field(default_factory=list)
@@ -180,162 +202,334 @@ RESTORE_ORDER = [
"workflow_trigger_logs",
]
DELETE_ORDER = [
"workflow_pause_reasons",
"workflow_node_execution_offload",
"workflow_trigger_logs",
"workflow_app_logs",
"workflow_node_executions",
"workflow_pauses",
"workflow_runs",
]
class WorkflowRunBundleArchiveMaintenance:
"""
Delete and restore V2 workflow-run archive bundles.
Delete accepts already-missing source rows only when every remaining row in the bundle scope is an unchanged
archive subset. It then removes only those verified primary keys, so unrelated or unarchived rows fail closed.
Non-dry-run delete and restore serialize on the existing archive catalog row before checking markers or changing
source rows.
Args:
dry_run: Validate and report counts without changing source rows or object-store markers.
strict_content_validation: Compare source-table content checksums against Parquet content before destructive
delete and after restore. Keep enabled for real maintenance.
stop_on_error: Stop batch processing after the first failed bundle.
strict_content_validation: Compare restored source-table content checksums against Parquet content. Delete
always validates that every remaining live row belongs to and matches the archive before removing it.
storage: Optional archive storage implementation. Tests may provide an in-memory implementation.
session_factory: Optional session factory. Each candidate is processed in its own transaction.
Batches stop at the first error so a returned cursor cannot pass an unhandled candidate.
"""
dry_run: bool
strict_content_validation: bool
stop_on_error: bool
storage: ArchiveStorage | None
session_factory: sessionmaker[Session]
def __init__(
self,
*,
dry_run: bool = False,
strict_content_validation: bool = True,
stop_on_error: bool = True,
storage: ArchiveStorage | None = None,
session_factory: sessionmaker[Session] | None = None,
) -> None:
self.dry_run = dry_run
self.strict_content_validation = strict_content_validation
self.stop_on_error = stop_on_error
self.storage = storage
self.session_factory = session_factory or sessionmaker(bind=db.engine, expire_on_commit=False)
def delete_batch(
self,
*,
tenant_ids: Sequence[str] | None,
start_date: datetime.datetime,
end_date: datetime.datetime,
target_year: int,
target_month: int,
after_catalog_id: str | None,
limit: int,
shard: str | None = None,
) -> BundleOperationSummary:
"""Validate and delete source rows for archived V2 bundles in the requested created_at window."""
"""Validate and delete one keyset page, optionally scoped to an exact archive shard."""
return self._process_batch(
operation="delete",
tenant_ids=tenant_ids,
start_date=start_date,
end_date=end_date,
target_year=target_year,
target_month=target_month,
after_catalog_id=after_catalog_id,
limit=limit,
shard=shard,
)
def restore_batch(
self,
*,
tenant_ids: Sequence[str] | None,
start_date: datetime.datetime,
end_date: datetime.datetime,
target_year: int,
target_month: int,
after_catalog_id: str | None,
limit: int,
) -> BundleOperationSummary:
"""Restore source rows for deleted V2 bundles in the requested created_at window."""
"""Restore source rows for a keyset page of deleted V2 bundles in one calendar month."""
return self._process_batch(
operation="restore",
tenant_ids=tenant_ids,
start_date=start_date,
end_date=end_date,
target_year=target_year,
target_month=target_month,
after_catalog_id=after_catalog_id,
limit=limit,
shard=None,
)
def validate_catalog_shards(
self,
*,
target_year: int,
target_month: int,
shard_total: int,
tenant_ids: Sequence[str] | None = None,
) -> None:
"""
Fail before a parallel delete when the requested closed-month scope contains a different shard layout.
A subset of the expected shards is valid because an archive shard may legitimately contain no bundles. Any
other shard name indicates a historical or mixed-layout month that must be handled by the serial delete path.
"""
if not 1 <= shard_total <= 16:
raise ValueError("shard_total must be between 1 and 16")
expected_shards = tuple(f"{index:02d}-of-{shard_total:02d}" for index in range(shard_total))
conditions = [
WorkflowRunArchiveBundle.year == target_year,
WorkflowRunArchiveBundle.month == target_month,
]
if tenant_ids is not None:
conditions.append(WorkflowRunArchiveBundle.tenant_id.in_(tenant_ids))
statement = (
select(WorkflowRunArchiveBundle.shard)
.where(
*conditions,
WorkflowRunArchiveBundle.shard.not_in(expected_shards),
)
.distinct()
.order_by(WorkflowRunArchiveBundle.shard.asc())
)
with self.session_factory() as session:
unexpected_shards = list(session.scalars(statement))
if unexpected_shards:
raise ValueError(
"archive catalog month contains unexpected shards for "
f"{shard_total}-way delete: {', '.join(unexpected_shards)}"
)
def _process_batch(
self,
*,
operation: str,
tenant_ids: Sequence[str] | None,
start_date: datetime.datetime,
end_date: datetime.datetime,
target_year: int,
target_month: int,
after_catalog_id: str | None,
limit: int,
shard: str | None,
) -> BundleOperationSummary:
start_time = time.time()
summary = BundleOperationSummary(operation=operation)
if tenant_ids is not None and not tenant_ids:
return summary
storage = self._get_archive_storage()
bundle_refs = self._list_bundle_refs(
storage,
operation=operation,
storage = self.storage or self._get_archive_storage()
catalog_entries = self._list_catalog_entries(
tenant_ids=tenant_ids,
start_date=start_date,
end_date=end_date,
target_year=target_year,
target_month=target_month,
after_catalog_id=after_catalog_id,
limit=limit,
shard=shard,
)
logger.info("Found %s V2 archive bundles for %s", len(bundle_refs), operation)
session_maker = sessionmaker(bind=db.engine, expire_on_commit=False)
for bundle_ref in bundle_refs:
with session_maker() as session:
if operation == "delete":
result = self._delete_bundle(session, storage, bundle_ref)
elif operation == "restore":
result = self._restore_bundle(session, storage, bundle_ref)
else:
raise ValueError(f"Unsupported operation: {operation}")
logger.info(
"Found %s V2 archive catalog candidates for %s: year=%s month=%s shard=%s after_catalog_id=%s",
len(catalog_entries),
operation,
target_year,
target_month,
shard,
after_catalog_id,
)
for catalog_entry in catalog_entries:
try:
bundle_ref = self._build_bundle_reference(storage, catalog_entry)
with self.session_factory() as session:
if operation == "delete":
result = self._delete_bundle(session, storage, bundle_ref)
elif operation == "restore":
result = self._restore_bundle(session, storage, bundle_ref)
else:
raise ValueError(f"Unsupported operation: {operation}")
except Exception as exc:
result = self._new_result_from_catalog_entry(catalog_entry)
result.error = str(exc)
logger.exception(
"Failed to prepare V2 archive bundle %s from catalog %s",
catalog_entry.bundle_id,
catalog_entry.catalog_id,
)
self._merge_result(summary, result)
if not result.success and self.stop_on_error:
if result.success:
if self.dry_run:
summary.preview_next_catalog_id = catalog_entry.catalog_id
else:
summary.next_catalog_id = catalog_entry.catalog_id
else:
logger.error("Stopping V2 bundle %s after failure: %s", operation, result.error)
break
summary.elapsed_time = time.time() - start_time
return summary
def _list_bundle_refs(
def _list_catalog_entries(
self,
storage: ArchiveStorage,
*,
operation: str,
tenant_ids: Sequence[str] | None,
start_date: datetime.datetime,
end_date: datetime.datetime,
target_year: int,
target_month: int,
after_catalog_id: str | None,
limit: int,
) -> list[BundleReference]:
start_date = self._to_naive_utc(start_date)
end_date = self._to_naive_utc(end_date)
manifest_keys = self._list_manifest_keys(storage, tenant_ids)
refs: list[BundleReference] = []
for manifest_key in manifest_keys:
manifest_data = self._get_checked_object(storage, manifest_key)
object_prefix = manifest_key.removesuffix(f"/{ARCHIVE_BUNDLE_MANIFEST_NAME}")
manifest = self._load_and_validate_manifest(manifest_data, object_prefix=object_prefix)
min_created_at = self._parse_manifest_datetime(manifest["min_created_at"])
max_created_at = self._parse_manifest_datetime(manifest["max_created_at"])
if max_created_at < start_date or min_created_at >= end_date:
continue
if tenant_ids and manifest["tenant_id"] not in tenant_ids:
continue
if operation == "delete" and self._is_deleted(storage, object_prefix):
continue
if operation == "restore" and not self._is_deleted(storage, object_prefix):
continue
refs.append(BundleReference(object_prefix=object_prefix, manifest_key=manifest_key, manifest=manifest))
shard: str | None = None,
) -> list[ArchiveBundleCatalogEntry]:
"""Read one bounded, stable-order candidate page from the database catalog."""
conditions = [
WorkflowRunArchiveBundle.year == target_year,
WorkflowRunArchiveBundle.month == target_month,
]
if tenant_ids:
conditions.append(WorkflowRunArchiveBundle.tenant_id.in_(tenant_ids))
if shard is not None:
conditions.append(WorkflowRunArchiveBundle.shard == shard)
if after_catalog_id:
conditions.append(WorkflowRunArchiveBundle.id > after_catalog_id)
refs.sort(
key=lambda ref: (
self._parse_manifest_datetime(ref.manifest["min_created_at"]),
ref.manifest["tenant_id"],
ref.manifest["bundle_id"],
)
statement = (
select(WorkflowRunArchiveBundle).where(*conditions).order_by(WorkflowRunArchiveBundle.id.asc()).limit(limit)
)
return refs[:limit]
with self.session_factory() as session:
if after_catalog_id:
cursor_bundle = session.get(WorkflowRunArchiveBundle, after_catalog_id)
self._validate_catalog_cursor_scope(
cursor_bundle,
tenant_ids=tenant_ids,
target_year=target_year,
target_month=target_month,
shard=shard,
)
bundles = list(session.scalars(statement))
return [
ArchiveBundleCatalogEntry(
catalog_id=bundle.id,
tenant_id=bundle.tenant_id,
year=bundle.year,
month=bundle.month,
shard=bundle.shard,
bundle_id=bundle.bundle_id,
workflow_run_count=bundle.workflow_run_count,
row_count=bundle.row_count,
archive_bytes=bundle.archive_bytes,
)
for bundle in bundles
]
@staticmethod
def _list_manifest_keys(storage: ArchiveStorage, tenant_ids: Sequence[str] | None) -> list[str]:
keys: list[str] = []
if tenant_ids:
prefixes = [
f"{_ARCHIVE_ROOT_PREFIX}tenant_prefix={tenant_id[0].lower()}/tenant_id={tenant_id}/"
for tenant_id in tenant_ids
]
else:
prefixes = [_ARCHIVE_ROOT_PREFIX]
for prefix in prefixes:
keys.extend(storage.list_objects(prefix))
return sorted(key for key in keys if key.endswith(f"/{ARCHIVE_BUNDLE_MANIFEST_NAME}"))
def _validate_catalog_cursor_scope(
cursor_bundle: WorkflowRunArchiveBundle | None,
*,
tenant_ids: Sequence[str] | None,
target_year: int,
target_month: int,
shard: str | None,
) -> None:
"""Reject a keyset cursor that cannot safely represent the requested catalog scope."""
if cursor_bundle is None:
raise ValueError("after_catalog_id does not exist in the workflow run archive bundle catalog")
if cursor_bundle.year != target_year or cursor_bundle.month != target_month:
raise ValueError("after_catalog_id is outside the requested archive month")
if tenant_ids is not None and cursor_bundle.tenant_id not in tenant_ids:
raise ValueError("after_catalog_id is outside the requested tenant scope")
if shard is not None and cursor_bundle.shard != shard:
raise ValueError("after_catalog_id is outside the requested archive shard")
def _build_bundle_reference(
self,
storage: ArchiveStorage,
catalog_entry: ArchiveBundleCatalogEntry,
) -> BundleReference:
"""Load and bind one manifest to the catalog row that selected it."""
object_prefix = self._catalog_object_prefix(catalog_entry)
manifest_key = f"{object_prefix}/{ARCHIVE_BUNDLE_MANIFEST_NAME}"
manifest_data = storage.get_object(manifest_key)
manifest = self._load_and_validate_manifest(manifest_data, object_prefix=object_prefix)
self._validate_manifest_catalog_identity(manifest, catalog_entry)
return BundleReference(
catalog=catalog_entry,
object_prefix=object_prefix,
manifest_key=manifest_key,
manifest_size_bytes=len(manifest_data),
manifest=manifest,
)
@staticmethod
def _catalog_object_prefix(catalog_entry: ArchiveBundleCatalogEntry) -> str:
"""Construct the immutable V2 bundle prefix from its database catalog identity."""
if not catalog_entry.tenant_id:
raise ValueError("archive catalog tenant_id must not be empty")
if not 1 <= catalog_entry.month <= 12:
raise ValueError(f"archive catalog month is invalid: {catalog_entry.month}")
return (
f"workflow-runs/v2/tenant_prefix={catalog_entry.tenant_id[0].lower()}/"
f"tenant_id={catalog_entry.tenant_id}/year={catalog_entry.year:04d}/"
f"month={catalog_entry.month:02d}/shard={catalog_entry.shard}/bundle={catalog_entry.bundle_id}"
)
@staticmethod
def _validate_manifest_catalog_identity(
manifest: BundleManifest,
catalog_entry: ArchiveBundleCatalogEntry,
) -> None:
"""Fail closed when the catalog locator and manifest identify different immutable bundles."""
expected_identity = (
catalog_entry.tenant_id,
catalog_entry.year,
catalog_entry.month,
catalog_entry.shard,
catalog_entry.bundle_id,
)
manifest_identity = (
manifest["tenant_id"],
manifest["year"],
manifest["month"],
manifest["shard"],
manifest["bundle_id"],
)
if manifest_identity != expected_identity:
raise ValueError(
f"archive manifest identity does not match catalog: expected={expected_identity}, "
f"actual={manifest_identity}"
)
if manifest["workflow_run_count"] != catalog_entry.workflow_run_count:
raise ValueError("archive manifest workflow_run_count does not match catalog")
manifest_row_count = sum(table["row_count"] for table in manifest["tables"].values())
if manifest_row_count != catalog_entry.row_count:
raise ValueError("archive manifest row_count does not match catalog")
def _delete_bundle(
self,
@@ -344,37 +538,61 @@ class WorkflowRunBundleArchiveMaintenance:
bundle_ref: BundleReference,
) -> BundleOperationResult:
start_time = time.time()
result = self._new_result(bundle_ref.manifest)
result = self._new_result(bundle_ref.manifest, bundle_ref.catalog.catalog_id)
try:
validation_start = time.time()
if not self.dry_run:
self._lock_catalog_entry(session, bundle_ref.catalog)
if self._is_restore_started(storage, bundle_ref.object_prefix):
raise ValueError("restore started marker exists; reconcile restore before delete")
deleted_marker_exists = self._is_deleted(storage, bundle_ref.object_prefix)
manifest, table_records, archive_bytes = self._validate_archive_object(storage, bundle_ref)
result.table_counts = self._manifest_table_counts(manifest)
result.archive_bytes = archive_bytes
self._lock_workflow_runs(session, manifest["run_ids"])
if self._is_delete_started(storage, bundle_ref.object_prefix) and self._live_counts_match(
session, manifest, expected_present=False
):
live_records = self._load_live_bundle_records(
session,
manifest,
table_records,
lock=not self.dry_run,
)
if deleted_marker_exists:
live_counts = {table_name: len(live_records[table_name]) for table_name in ARCHIVED_TABLES}
if any(live_counts.values()):
raise ValueError(f"Live rows exist for bundle with deleted marker: {live_counts}")
if not self.dry_run:
self._delete_marker(storage, bundle_ref.object_prefix, ARCHIVE_BUNDLE_DELETE_STARTED_MARKER_NAME)
self._delete_marker(storage, bundle_ref.object_prefix, ARCHIVE_BUNDLE_RESTORED_MARKER_NAME)
result.validation_time = time.time() - validation_start
result.success = True
result.elapsed_time = time.time() - start_time
return result
self._validate_live_archive_subset(manifest, table_records, live_records)
result.validation_time = time.time() - validation_start
if not any(live_records[table_name] for table_name in ARCHIVED_TABLES):
if not self.dry_run:
self._mark_deleted(storage, bundle_ref.object_prefix)
self._delete_marker(storage, bundle_ref.object_prefix, ARCHIVE_BUNDLE_DELETE_STARTED_MARKER_NAME)
self._delete_marker(storage, bundle_ref.object_prefix, ARCHIVE_BUNDLE_RESTORED_MARKER_NAME)
result.success = True
result.elapsed_time = time.time() - start_time
return result
self._validate_live_counts(session, manifest, expected_present=True)
if self.strict_content_validation:
self._validate_live_content(session, table_records)
result.validation_time = time.time() - validation_start
if not self.dry_run:
self._put_marker(storage, bundle_ref.object_prefix, ARCHIVE_BUNDLE_DELETE_STARTED_MARKER_NAME)
deleted_counts = self._delete_bundle_rows(session, table_records)
if deleted_counts != result.table_counts:
expected_deleted_counts = {table_name: len(live_records[table_name]) for table_name in ARCHIVED_TABLES}
deleted_counts = self._delete_bundle_rows(session, live_records)
if deleted_counts != expected_deleted_counts:
raise ValueError(
f"Deleted row count mismatch: expected={result.table_counts}, actual={deleted_counts}"
f"Deleted row count mismatch: expected={expected_deleted_counts}, actual={deleted_counts}"
)
self._validate_live_counts(session, manifest, expected_present=False)
remaining_records = self._load_live_bundle_records(session, manifest, table_records, lock=True)
remaining_counts = {table_name: len(remaining_records[table_name]) for table_name in ARCHIVED_TABLES}
if any(remaining_counts.values()):
raise ValueError(f"Live rows remain after bundle delete: {remaining_counts}")
session.commit()
self._mark_deleted(storage, bundle_ref.object_prefix)
self._delete_marker(storage, bundle_ref.object_prefix, ARCHIVE_BUNDLE_DELETE_STARTED_MARKER_NAME)
@@ -394,9 +612,25 @@ class WorkflowRunBundleArchiveMaintenance:
bundle_ref: BundleReference,
) -> BundleOperationResult:
start_time = time.time()
result = self._new_result(bundle_ref.manifest)
result = self._new_result(bundle_ref.manifest, bundle_ref.catalog.catalog_id)
try:
validation_start = time.time()
if not self.dry_run:
self._lock_catalog_entry(session, bundle_ref.catalog)
if not self._is_deleted(storage, bundle_ref.object_prefix):
# A committed delete may be interrupted before `.deleted` is written. Do not let restore advance its
# cursor over that state: retry delete to reconcile it, or investigate source-row drift first.
if self._is_delete_started(storage, bundle_ref.object_prefix):
raise ValueError("delete started marker exists without a deleted marker; reconcile delete first")
restore_started = self._is_restore_started(storage, bundle_ref.object_prefix)
self._validate_live_counts(session, bundle_ref.manifest, expected_present=True)
result.validation_time = time.time() - validation_start
if restore_started and not self.dry_run:
self._mark_restored(storage, bundle_ref.object_prefix)
result.success = True
result.elapsed_time = time.time() - start_time
return result
manifest, table_records, archive_bytes = self._validate_archive_object(storage, bundle_ref)
result.table_counts = self._manifest_table_counts(manifest)
result.archive_bytes = archive_bytes
@@ -408,6 +642,7 @@ class WorkflowRunBundleArchiveMaintenance:
if not self.dry_run:
self._mark_restored(storage, bundle_ref.object_prefix)
result.success = True
result.elapsed_time = time.time() - start_time
return result
self._validate_live_counts(session, manifest, expected_present=False)
@@ -432,13 +667,40 @@ class WorkflowRunBundleArchiveMaintenance:
return result
@staticmethod
def _new_result(manifest: BundleManifest) -> BundleOperationResult:
def _new_result(manifest: BundleManifest, catalog_id: str) -> BundleOperationResult:
return BundleOperationResult(
catalog_id=catalog_id,
bundle_id=manifest["bundle_id"],
tenant_id=manifest["tenant_id"],
object_prefix=manifest["object_prefix"],
)
@staticmethod
def _new_result_from_catalog_entry(catalog_entry: ArchiveBundleCatalogEntry) -> BundleOperationResult:
try:
object_prefix = WorkflowRunBundleArchiveMaintenance._catalog_object_prefix(catalog_entry)
except ValueError:
object_prefix = "<invalid archive catalog identity>"
return BundleOperationResult(
catalog_id=catalog_entry.catalog_id,
bundle_id=catalog_entry.bundle_id,
tenant_id=catalog_entry.tenant_id,
object_prefix=object_prefix,
)
@staticmethod
def _lock_catalog_entry(session: Session, catalog_entry: ArchiveBundleCatalogEntry) -> None:
locked_catalog_id = session.scalar(
select(WorkflowRunArchiveBundle.id)
.where(
WorkflowRunArchiveBundle.id == catalog_entry.catalog_id,
WorkflowRunArchiveBundle.tenant_id == catalog_entry.tenant_id,
)
.with_for_update()
)
if locked_catalog_id is None:
raise ValueError("archive catalog row disappeared before bundle maintenance")
def _validate_archive_object(
self,
storage: ArchiveStorage,
@@ -446,7 +708,7 @@ class WorkflowRunBundleArchiveMaintenance:
) -> tuple[BundleManifest, dict[str, list[dict[str, Any]]], int]:
manifest = bundle_ref.manifest
table_records: dict[str, list[dict[str, Any]]] = {}
total_size = len(storage.get_object(bundle_ref.manifest_key))
total_size = bundle_ref.manifest_size_bytes
for table_name in ARCHIVED_TABLES:
info = manifest["tables"][table_name]
payload = self._get_checked_object(storage, info["object_key"])
@@ -469,12 +731,14 @@ class WorkflowRunBundleArchiveMaintenance:
f"expected={info['row_count']}, actual={len(records)}"
)
table_records[table_name] = records
if total_size != bundle_ref.catalog.archive_bytes:
raise ValueError(
f"Archive object total size mismatch: expected={bundle_ref.catalog.archive_bytes}, actual={total_size}"
)
return manifest, table_records, total_size
@staticmethod
def _get_checked_object(storage: ArchiveStorage, object_key: str) -> bytes:
if not storage.object_exists(object_key):
raise FileNotFoundError(f"Archive object not found: {object_key}")
return storage.get_object(object_key)
@staticmethod
@@ -539,6 +803,110 @@ class WorkflowRunBundleArchiveMaintenance:
table = pq.read_table(io.BytesIO(payload))
return table.to_pylist()
def _load_live_bundle_records(
self,
session: Session,
manifest: BundleManifest,
table_records: dict[str, list[dict[str, Any]]],
*,
lock: bool,
) -> dict[str, list[dict[str, Any]]]:
"""Load the complete live scope for a bundle, including archived rows whose relationship fields drifted."""
run_ids = manifest["run_ids"]
archive_ids = {
table_name: [str(record["id"]) for record in table_records[table_name]] for table_name in ARCHIVED_TABLES
}
live_node_ids = self._select_ids_by_run_ids(session, WorkflowNodeExecutionModel, run_ids)
live_pause_ids = self._select_ids_by_run_ids(session, WorkflowPause, run_ids)
node_ids = sorted(set(archive_ids["workflow_node_executions"]) | set(live_node_ids))
pause_ids = sorted(set(archive_ids["workflow_pauses"]) | set(live_pause_ids))
def load_scope(
table_name: str,
model: Any,
scope_column: Any,
scope_ids: Sequence[str],
) -> list[dict[str, Any]]:
return self._merge_records_by_id(
self._load_records_by_column(session, model, scope_column, scope_ids, lock=lock),
self._load_records_by_column(session, model, model.id, archive_ids[table_name], lock=lock),
)
return {
"workflow_pause_reasons": load_scope(
"workflow_pause_reasons", WorkflowPauseReason, WorkflowPauseReason.pause_id, pause_ids
),
"workflow_node_execution_offload": load_scope(
"workflow_node_execution_offload",
WorkflowNodeExecutionOffload,
WorkflowNodeExecutionOffload.node_execution_id,
node_ids,
),
"workflow_trigger_logs": load_scope(
"workflow_trigger_logs", WorkflowTriggerLog, WorkflowTriggerLog.workflow_run_id, run_ids
),
"workflow_app_logs": load_scope(
"workflow_app_logs", WorkflowAppLog, WorkflowAppLog.workflow_run_id, run_ids
),
"workflow_node_executions": load_scope(
"workflow_node_executions",
WorkflowNodeExecutionModel,
WorkflowNodeExecutionModel.workflow_run_id,
run_ids,
),
"workflow_pauses": load_scope("workflow_pauses", WorkflowPause, WorkflowPause.workflow_run_id, run_ids),
"workflow_runs": self._load_records_by_column(
session,
WorkflowRun,
WorkflowRun.id,
sorted(set(run_ids) | set(archive_ids["workflow_runs"])),
lock=lock,
),
}
@classmethod
def _validate_live_archive_subset(
cls,
manifest: BundleManifest,
table_records: dict[str, list[dict[str, Any]]],
live_records: dict[str, list[dict[str, Any]]],
) -> None:
"""Require every live row in the bundle scope to exist unchanged in the validated archive."""
manifest_run_ids = {str(run_id) for run_id in manifest["run_ids"]}
if len(manifest_run_ids) != len(manifest["run_ids"]):
raise ValueError("archive manifest contains duplicate workflow run IDs")
archive_records_by_id: dict[str, dict[str, dict[str, Any]]] = {}
for table_name in ARCHIVED_TABLES:
records_by_id = {str(record["id"]): record for record in table_records[table_name]}
if len(records_by_id) != len(table_records[table_name]):
raise ValueError(f"archive contains duplicate row IDs for {table_name}")
archive_records_by_id[table_name] = records_by_id
if set(archive_records_by_id["workflow_runs"]) != manifest_run_ids:
raise ValueError("archive workflow run IDs do not match manifest run_ids")
for table_name in ARCHIVED_TABLES:
live_ids = [str(record["id"]) for record in live_records[table_name]]
if len(set(live_ids)) != len(live_ids):
raise ValueError(f"live scope contains duplicate row IDs for {table_name}")
archive_by_id = archive_records_by_id[table_name]
extra_ids = sorted(set(live_ids) - set(archive_by_id))
if extra_ids:
raise ValueError(
f"Live bundle scope contains rows missing from archive for {table_name}: {extra_ids[:10]}"
)
archive_subset = [archive_by_id[row_id] for row_id in live_ids]
live_checksum = cls._records_checksum(live_records[table_name])
archive_checksum = cls._records_checksum(archive_subset)
if live_checksum != archive_checksum:
raise ValueError(
f"Live/archive subset content checksum mismatch for {table_name}: "
f"expected={archive_checksum}, actual={live_checksum}"
)
def _validate_live_counts(
self,
session: Session,
@@ -617,26 +985,13 @@ class WorkflowRunBundleArchiveMaintenance:
def _delete_bundle_rows(
self,
session: Session,
table_records: dict[str, list[dict[str, Any]]],
live_records: dict[str, list[dict[str, Any]]],
) -> dict[str, int]:
run_ids = [str(record["id"]) for record in table_records["workflow_runs"]]
node_ids = [str(record["id"]) for record in table_records["workflow_node_executions"]]
pause_ids = [str(record["id"]) for record in table_records["workflow_pauses"]]
deleted_counts = dict.fromkeys(ARCHIVED_TABLES, 0)
deleted_counts["workflow_pause_reasons"] = self._delete_by_column(
session, WorkflowPauseReason, WorkflowPauseReason.pause_id, pause_ids
)
deleted_counts["workflow_node_execution_offload"] = self._delete_by_column(
session, WorkflowNodeExecutionOffload, WorkflowNodeExecutionOffload.node_execution_id, node_ids
)
deleted_counts["workflow_trigger_logs"] = self._delete_by_run_ids(session, WorkflowTriggerLog, run_ids)
deleted_counts["workflow_app_logs"] = self._delete_by_run_ids(session, WorkflowAppLog, run_ids)
deleted_counts["workflow_node_executions"] = self._delete_by_run_ids(
session, WorkflowNodeExecutionModel, run_ids
)
deleted_counts["workflow_pauses"] = self._delete_by_run_ids(session, WorkflowPause, run_ids)
deleted_counts["workflow_runs"] = self._delete_by_run_ids(session, WorkflowRun, run_ids)
for table_name in DELETE_ORDER:
model = TABLE_MODELS[table_name]
row_ids = [str(record["id"]) for record in live_records[table_name]]
deleted_counts[table_name] = self._delete_by_column(session, model, model.id, row_ids)
return deleted_counts
def _restore_bundle_rows(
@@ -708,11 +1063,6 @@ class WorkflowRunBundleArchiveMaintenance:
payload = json.dumps(normalized, sort_keys=True, default=str, ensure_ascii=False, separators=(",", ":"))
return ArchiveStorage.compute_checksum(payload.encode("utf-8"))
@staticmethod
def _lock_workflow_runs(session: Session, run_ids: Sequence[str]) -> None:
for chunk in WorkflowRunBundleArchiveMaintenance._chunks(run_ids, _CHUNK_SIZE):
list(session.scalars(select(WorkflowRun.id).where(WorkflowRun.id.in_(chunk)).with_for_update()))
@staticmethod
def _select_ids_by_run_ids(
session: Session,
@@ -757,8 +1107,16 @@ class WorkflowRunBundleArchiveMaintenance:
session: Session,
model: Any,
run_ids: Sequence[str],
*,
lock: bool = False,
) -> list[dict[str, Any]]:
return self._load_records_by_column(session, model, self._run_id_column(model), run_ids)
return self._load_records_by_column(
session,
model,
self._run_id_column(model),
run_ids,
lock=lock,
)
def _load_records_by_column(
self,
@@ -766,23 +1124,23 @@ class WorkflowRunBundleArchiveMaintenance:
model: Any,
column: Any,
values: Sequence[str],
*,
lock: bool = False,
) -> list[dict[str, Any]]:
if not values:
return []
rows: list[Any] = []
for chunk in self._chunks(values, _CHUNK_SIZE):
rows.extend(session.scalars(select(model).where(column.in_(chunk))))
statement = select(model).where(column.in_(chunk)).order_by(model.id.asc())
if lock:
statement = statement.with_for_update()
rows.extend(session.scalars(statement))
return [self._row_to_dict(row) for row in rows]
@staticmethod
def _delete_by_run_ids(
session: Session,
model: Any,
run_ids: Sequence[str],
) -> int:
return WorkflowRunBundleArchiveMaintenance._delete_by_column(
session, model, WorkflowRunBundleArchiveMaintenance._run_id_column(model), run_ids
)
def _merge_records_by_id(*record_groups: list[dict[str, Any]]) -> list[dict[str, Any]]:
records_by_id = {str(record["id"]): record for record_group in record_groups for record in record_group}
return [records_by_id[row_id] for row_id in sorted(records_by_id)]
@staticmethod
def _run_id_column(model: Any) -> Any:
@@ -813,6 +1171,10 @@ class WorkflowRunBundleArchiveMaintenance:
def _is_delete_started(storage: ArchiveStorage, object_prefix: str) -> bool:
return storage.object_exists(f"{object_prefix}/{ARCHIVE_BUNDLE_DELETE_STARTED_MARKER_NAME}")
@staticmethod
def _is_restore_started(storage: ArchiveStorage, object_prefix: str) -> bool:
return storage.object_exists(f"{object_prefix}/{ARCHIVE_BUNDLE_RESTORE_STARTED_MARKER_NAME}")
@staticmethod
def _mark_deleted(storage: ArchiveStorage, object_prefix: str) -> None:
WorkflowRunBundleArchiveMaintenance._put_marker(storage, object_prefix, ARCHIVE_BUNDLE_DELETED_MARKER_NAME)
@@ -820,10 +1182,13 @@ class WorkflowRunBundleArchiveMaintenance:
@staticmethod
def _mark_restored(storage: ArchiveStorage, object_prefix: str) -> None:
WorkflowRunBundleArchiveMaintenance._delete_marker(storage, object_prefix, ARCHIVE_BUNDLE_DELETED_MARKER_NAME)
WorkflowRunBundleArchiveMaintenance._put_marker(storage, object_prefix, ARCHIVE_BUNDLE_RESTORED_MARKER_NAME)
WorkflowRunBundleArchiveMaintenance._delete_marker(
storage, object_prefix, ARCHIVE_BUNDLE_DELETE_STARTED_MARKER_NAME
)
WorkflowRunBundleArchiveMaintenance._delete_marker(
storage, object_prefix, ARCHIVE_BUNDLE_RESTORE_STARTED_MARKER_NAME
)
WorkflowRunBundleArchiveMaintenance._put_marker(storage, object_prefix, ARCHIVE_BUNDLE_RESTORED_MARKER_NAME)
@staticmethod
def _put_marker(storage: ArchiveStorage, object_prefix: str, marker_name: str) -> None:
@@ -836,16 +1201,6 @@ class WorkflowRunBundleArchiveMaintenance:
if storage.object_exists(marker_key):
storage.delete_object(marker_key)
@staticmethod
def _parse_manifest_datetime(value: str) -> datetime.datetime:
return WorkflowRunBundleArchiveMaintenance._to_naive_utc(datetime.datetime.fromisoformat(value))
@staticmethod
def _to_naive_utc(value: datetime.datetime) -> datetime.datetime:
if value.tzinfo is None:
return value
return value.astimezone(datetime.UTC).replace(tzinfo=None)
@staticmethod
def _chunks(values: Sequence[Any], size: int) -> list[Sequence[Any]]:
return [values[index : index + size] for index in range(0, len(values), size)]
+221
View File
@@ -0,0 +1,221 @@
"""Account-level Step-by-step Tour persistence."""
from datetime import datetime
from typing import NotRequired, TypedDict
from sqlalchemy import select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session, scoped_session
from configs import dify_config
from libs.datetime_utils import ensure_naive_utc
from models.account import Account
from models.onboarding import AccountStepByStepTourState
STEP_BY_STEP_TOUR_TASK_IDS = frozenset(("home", "studio", "knowledge", "integration"))
class StepByStepTourStateResponse(TypedDict):
first_workspace_id: str | None
skipped: bool
completed_task_ids: list[str]
manually_enabled_workspace_ids: list[str]
manually_disabled_workspace_ids: list[str]
updated_at: datetime | None
class StepByStepTourPatch(TypedDict):
action: str
task_id: NotRequired[str | None]
class StepByStepTourService:
"""Coordinate persisted tour state with account eligibility rules."""
@classmethod
def get_state(
cls,
*,
account: Account,
current_tenant_id: str,
session: Session | scoped_session,
) -> StepByStepTourStateResponse:
eligible = cls.is_eligible(account)
state = cls._get_state(account.id, session=session)
if eligible:
state = cls._ensure_state(account.id, session=session, state=state)
if state.first_workspace_id is None:
state.first_workspace_id = current_tenant_id
session.commit()
session.refresh(state)
return cls._build_response(state=state)
@classmethod
def patch_state(
cls,
*,
account: Account,
current_tenant_id: str,
patch: StepByStepTourPatch,
session: Session | scoped_session,
) -> StepByStepTourStateResponse:
state = cls._ensure_state(account.id, session=session, state=None)
cls._apply_action(
state=state,
action=patch["action"],
task_id=patch.get("task_id"),
current_tenant_id=current_tenant_id,
)
session.commit()
session.refresh(state)
return cls._build_response(state=state)
@classmethod
def is_eligible(cls, account: Account) -> bool:
if not dify_config.ENABLE_STEP_BY_STEP_TOUR:
return False
rollout_started_at = dify_config.STEP_BY_STEP_TOUR_ROLLOUT_STARTED_AT
if rollout_started_at is None:
return False
account_started_at = account.initialized_at or account.created_at
if account_started_at is None:
return False
return ensure_naive_utc(account_started_at) >= ensure_naive_utc(rollout_started_at)
@classmethod
def _get_state(
cls,
account_id: str,
*,
session: Session | scoped_session,
) -> AccountStepByStepTourState | None:
stmt = select(AccountStepByStepTourState).where(AccountStepByStepTourState.account_id == account_id).limit(1)
return session.execute(stmt).scalar_one_or_none()
@classmethod
def _ensure_state(
cls,
account_id: str,
*,
session: Session | scoped_session,
state: AccountStepByStepTourState | None,
) -> AccountStepByStepTourState:
if state is None:
state = cls._get_state(account_id, session=session)
if state is not None:
return state
state = AccountStepByStepTourState(account_id=account_id)
session.add(state)
try:
session.flush()
except IntegrityError:
# Another tab/device can create the account row between our read and insert.
session.rollback()
state = cls._get_state(account_id, session=session)
if state is None:
raise
return state
@classmethod
def _apply_action(
cls,
*,
state: AccountStepByStepTourState,
action: str,
task_id: str | None,
current_tenant_id: str,
) -> None:
match action:
case "skip":
state.skipped = True
state.manually_enabled_workspace_ids = cls._remove_id(
state.manually_enabled_workspace_ids,
current_tenant_id,
)
case "complete_task":
if task_id is None:
raise ValueError("task_id is required")
cls._validate_task_id(task_id)
state.completed_task_ids = cls._add_id(state.completed_task_ids, task_id)
case "uncomplete_task":
if task_id is None:
raise ValueError("task_id is required")
cls._validate_task_id(task_id)
state.completed_task_ids = cls._remove_id(state.completed_task_ids, task_id)
case "enable_current_workspace":
state.skipped = False
state.manually_enabled_workspace_ids = cls._add_id(
state.manually_enabled_workspace_ids,
current_tenant_id,
)
state.manually_disabled_workspace_ids = cls._remove_id(
state.manually_disabled_workspace_ids,
current_tenant_id,
)
case "disable_current_workspace":
state.manually_enabled_workspace_ids = cls._remove_id(
state.manually_enabled_workspace_ids,
current_tenant_id,
)
state.manually_disabled_workspace_ids = cls._add_id(
state.manually_disabled_workspace_ids,
current_tenant_id,
)
case _:
raise ValueError(f"Unsupported action: {action}")
@classmethod
def _build_response(
cls,
*,
state: AccountStepByStepTourState | None,
) -> StepByStepTourStateResponse:
if state is None:
return {
"first_workspace_id": None,
"skipped": False,
"completed_task_ids": [],
"manually_enabled_workspace_ids": [],
"manually_disabled_workspace_ids": [],
"updated_at": None,
}
return {
"first_workspace_id": state.first_workspace_id,
"skipped": state.skipped,
"completed_task_ids": cls._normalize_ids(state.completed_task_ids),
"manually_enabled_workspace_ids": cls._normalize_ids(state.manually_enabled_workspace_ids),
"manually_disabled_workspace_ids": cls._normalize_ids(state.manually_disabled_workspace_ids),
"updated_at": state.updated_at,
}
@staticmethod
def _validate_task_id(task_id: str) -> None:
if task_id not in STEP_BY_STEP_TOUR_TASK_IDS:
raise ValueError(f"Unsupported task_id: {task_id}")
@classmethod
def _add_id(cls, values: list[str], value: str) -> list[str]:
normalized = cls._normalize_ids(values)
if value in normalized:
return normalized
return [*normalized, value]
@classmethod
def _remove_id(cls, values: list[str], value: str) -> list[str]:
return [item for item in cls._normalize_ids(values) if item != value]
@staticmethod
def _normalize_ids(values: list[str]) -> list[str]:
normalized: list[str] = []
for value in values:
if value not in normalized:
normalized.append(value)
return normalized
+161 -14
View File
@@ -8,6 +8,7 @@ import uuid
from collections.abc import Mapping
from typing import Any, override
from socketio.exceptions import TimeoutError as SocketIOTimeoutError # type: ignore[reportMissingTypeStubs]
from sqlalchemy import select
from sqlalchemy.orm import Session
@@ -18,6 +19,7 @@ from repositories.workflow_collaboration_repository import WorkflowCollaboration
logger = logging.getLogger(__name__)
SERVER_HEARTBEAT_INTERVAL_SECONDS = 30
SYNC_REQUEST_TIMEOUT_SECONDS = 15
_PROCESS_SERVER_ID: str | None = None
_PROCESS_SERVER_ID_PID: int | None = None
@@ -38,7 +40,8 @@ class WorkflowCollaborationService:
Socket.IO rooms are process-local unless backed by a message queue, while online users and leader election live in
Redis. Each websocket worker writes a small heartbeat keyed by `server_id`; session rows store their owner so a
worker can distinguish a live remote sid from a stale sid left behind by a dead worker.
worker can distinguish a live remote sid from a stale sid left behind by a dead worker. Visibility events are
ordered per socket, and sync requests wait for the selected saver to acknowledge its draft write.
"""
_heartbeat_started: bool
@@ -128,6 +131,9 @@ class WorkflowCollaborationService:
"sid": sid,
"connected_at": int(time.time()),
"server_id": self.server_id,
# Joins are assumed visible; hidden tabs re-report via a graph_view_state
# event right after receiving the post-join "status" emit.
"graph_active": True,
}
self._repository.set_session_info(workflow_id, session_info)
@@ -158,7 +164,8 @@ class WorkflowCollaborationService:
self.handle_leader_disconnect(workflow_id, sid)
self.broadcast_online_users(workflow_id)
def relay_collaboration_event(self, sid: str, data: Mapping[str, object]) -> tuple[dict[str, str], int]:
def relay_collaboration_event(self, sid: str, data: Mapping[str, object]) -> tuple[dict[str, object], int]:
"""Route collaboration control events, directing save and graph-resync requests to the active leader."""
mapping = self._repository.get_sid_mapping(sid)
if not mapping:
return {"msg": "unauthorized"}, 401
@@ -174,11 +181,54 @@ class WorkflowCollaborationService:
if not event_type:
return {"msg": "invalid event type"}, 400
if event_type == "graph_view_state":
if not isinstance(event_data, Mapping):
return {"msg": "invalid graph_view_state"}, 400
graph_active = event_data.get("graphActive")
sequence = event_data.get("sequence")
if not isinstance(graph_active, bool):
return {"msg": "invalid graph_view_state"}, 400
if not isinstance(sequence, int) or isinstance(sequence, bool) or sequence < 0:
return {"msg": "invalid graph_view_state"}, 400
# Sequence the visibility write together with its leader side effect. Otherwise a newer
# visible event can land after the Lua write but before an older hidden handler demotes.
with self._repository.graph_view_state_lock(workflow_id):
applied = self._repository.update_session_graph_active(workflow_id, sid, graph_active, sequence)
if not applied:
return {"msg": "graph_view_state_ignored"}, 200
# Write the flag before re-electing so tier-1 selection already excludes the
# session that just went hidden (including one _ensure_leader just promoted).
if not graph_active:
self._demote_leader_if_hidden(workflow_id, sid)
return {"msg": "graph_view_state_updated"}, 200
if event_type == "sync_request":
if not isinstance(event_data, Mapping):
return {"msg": "invalid sync_request"}, 400
request_id = event_data.get("requestId")
if not isinstance(request_id, str) or not request_id.strip():
return {"msg": "invalid sync_request"}, 400
leader_sid = self._repository.get_current_leader(workflow_id)
target_sid: str | None
if leader_sid and self.is_session_active(workflow_id, leader_sid):
target_sid = leader_sid
if self._is_session_graph_active(workflow_id, leader_sid):
target_sid = leader_sid
else:
# The leader is connected but its tab is hidden: its canvas is frozen
# (rAF paused), so saving through it would persist stale data. Hand
# leadership to a visible session — or, if every tab is hidden, to the
# requester, whose canvas at least contains its own edits.
replacement = self._select_graph_leader(workflow_id, preferred_sid=sid)
if replacement and replacement != leader_sid:
self._repository.set_leader(workflow_id, replacement)
self.broadcast_leader_change(workflow_id, replacement)
target_sid = replacement
else:
target_sid = leader_sid
else:
if leader_sid:
self._repository.delete_leader(workflow_id)
@@ -188,15 +238,73 @@ class WorkflowCollaborationService:
self.broadcast_leader_change(workflow_id, target_sid)
if not target_sid:
return {"msg": "no_active_leader"}, 200
return {"msg": "no_active_leader", "requestId": request_id}, 503
target_data = dict(event_data)
target_data["requestId"] = request_id
try:
result = self._socketio.call(
"collaboration_update",
{"type": event_type, "userId": user_id, "data": target_data, "timestamp": timestamp},
to=target_sid,
timeout=SYNC_REQUEST_TIMEOUT_SECONDS,
)
except SocketIOTimeoutError:
logger.warning(
"Workflow collaboration sync request timed out: workflow_id=%s requester_sid=%s target_sid=%s",
workflow_id,
sid,
target_sid,
)
return {"msg": "sync_request_timeout", "requestId": request_id}, 504
if not isinstance(result, Mapping) or result.get("success") is not True:
response: dict[str, object] = {
"msg": "workflow_sync_failed",
"requestId": request_id,
"success": False,
}
if isinstance(result, Mapping) and isinstance(result.get("error"), str):
response["error"] = result["error"]
return response, 502
workflow_hash = result.get("hash")
updated_at = result.get("updatedAt")
if not isinstance(workflow_hash, str) or not workflow_hash:
return {"msg": "invalid_sync_response", "requestId": request_id}, 502
if not isinstance(updated_at, int) or isinstance(updated_at, bool):
return {"msg": "invalid_sync_response", "requestId": request_id}, 502
return {
"msg": "workflow_synced",
"requestId": request_id,
"success": True,
"hash": workflow_hash,
"updatedAt": updated_at,
}, 200
if event_type == "graph_resync_request":
leader_sid = self._repository.get_current_leader(workflow_id)
resync_target_sid: str | None
if leader_sid and self.is_session_active(workflow_id, leader_sid):
resync_target_sid = leader_sid
else:
if leader_sid:
self._repository.delete_leader(workflow_id)
resync_target_sid = self._select_graph_leader(workflow_id, preferred_sid=sid)
if resync_target_sid:
self._repository.set_leader(workflow_id, resync_target_sid)
self.broadcast_leader_change(workflow_id, resync_target_sid)
if not resync_target_sid:
return {"msg": "no_active_leader"}, 503
self._socketio.emit(
"collaboration_update",
{"type": event_type, "userId": user_id, "data": event_data, "timestamp": timestamp},
room=target_sid,
to=resync_target_sid,
)
return {"msg": "sync_request_forwarded"}, 200
return {"msg": "graph_resync_request_forwarded"}, 200
self._socketio.emit(
"collaboration_update",
@@ -329,17 +437,56 @@ class WorkflowCollaborationService:
self._repository.set_leader(workflow_id, sid)
self.broadcast_leader_change(workflow_id, sid)
def _select_graph_leader(self, workflow_id: str, preferred_sid: str | None = None) -> str | None:
session_sids = [
session["sid"]
def _select_graph_leader(
self,
workflow_id: str,
preferred_sid: str | None = None,
*,
require_graph_active: bool = False,
) -> str | None:
"""Pick a leader, preferring sessions whose canvas tab is visible.
Hidden tabs freeze rAF-driven CRDT->canvas application, so a visible session is
always the freshest saver. When every tab is hidden the room must still keep a
leader (followers drop sync_requests unless they believe they are the leader),
so tier 2 falls back to any active session unless require_graph_active is set.
"""
active_sessions = [
session
for session in self._repository.list_sessions(workflow_id)
if session.get("graph_active", True) and self.is_session_active(workflow_id, session["sid"])
if self.is_session_active(workflow_id, session["sid"])
]
if not session_sids:
visible_sids = [session["sid"] for session in active_sessions if session.get("graph_active", True)]
candidate_sids = visible_sids
if not candidate_sids and not require_graph_active:
candidate_sids = [session["sid"] for session in active_sessions]
if not candidate_sids:
return None
if preferred_sid and preferred_sid in session_sids:
if preferred_sid and preferred_sid in candidate_sids:
return preferred_sid
return session_sids[0]
return candidate_sids[0]
def _is_session_graph_active(self, workflow_id: str, sid: str) -> bool:
"""Default to True on missing/unreadable session data so read failures never churn leadership."""
session_info = self._repository.get_session_info(workflow_id, sid)
if session_info is None:
return True
return bool(session_info.get("graph_active", True))
def _demote_leader_if_hidden(self, workflow_id: str, hidden_sid: str) -> None:
current_leader = self._repository.get_current_leader(workflow_id)
if current_leader != hidden_sid:
return
new_leader = self._select_graph_leader(workflow_id, require_graph_active=True)
# No visible session: keep the hidden leader rather than leaving the room leaderless.
if not new_leader or new_leader == hidden_sid:
return
self._repository.set_leader(workflow_id, new_leader)
self.broadcast_leader_change(workflow_id, new_leader)
def is_session_active(self, workflow_id: str, sid: str) -> bool:
if not sid:
+11 -10
View File
@@ -823,6 +823,8 @@ _FILENAME_TRANS_TABLE = _make_filename_trans_table()
class DraftVariableSaver:
"""Persist draft outputs under the tenant that owns the app or pipeline."""
# _DUMMY_OUTPUT_IDENTITY is a placeholder output for workflow nodes.
# Its sole possible value is `None`.
#
@@ -842,6 +844,10 @@ class DraftVariableSaver:
# Database session used for persisting draft variables.
_session: Session
# Resource owner tenant. An account's current tenant may be unset or point elsewhere
# when draft variables are persisted by an asynchronous workflow execution.
_tenant_id: str
# The application ID associated with the draft variables.
# This should match the `Workflow.app_id` of the workflow to which the current node belongs.
_app_id: str
@@ -867,6 +873,7 @@ class DraftVariableSaver:
def __init__(
self,
session: Session,
tenant_id: str,
app_id: str,
node_id: str,
node_type: NodeType,
@@ -878,6 +885,7 @@ class DraftVariableSaver:
# WorkflowNodeExecutionModel/WorkflowNodeExecution, not their `node_execution_id`
# field. These are distinct database fields with different purposes.
self._session = session
self._tenant_id = tenant_id
self._app_id = app_id
self._node_id = node_id
self._node_type = node_type
@@ -885,12 +893,6 @@ class DraftVariableSaver:
self._user = user
self._enclosing_node_id = enclosing_node_id
def _resolve_app_tenant_id(self) -> str:
tenant_id = self._session.scalar(select(App.tenant_id).where(App.id == self._app_id))
if not tenant_id:
raise ValueError(f"Unable to resolve tenant_id for app {self._app_id}")
return tenant_id
def _create_dummy_output_variable(self):
return WorkflowDraftVariable.new_node_variable(
app_id=self._app_id,
@@ -949,11 +951,10 @@ class DraftVariableSaver:
if name == SystemVariableKey.FILES:
# Here we know the type of variable must be `array[file]`, we
# just rebuild files from the serialized payload.
tenant_id = self._resolve_app_tenant_id()
files = [
build_file_from_stored_mapping(
file_mapping=v,
tenant_id=tenant_id,
tenant_id=self._tenant_id,
)
for v in value
]
@@ -1096,8 +1097,8 @@ class DraftVariableSaver:
content=original_content_serialized.encode(),
mimetype=content_type,
user=self._user,
tenant_id=self._tenant_id,
)
assert self._user.current_tenant_id
# Create WorkflowDraftVariableFile record
variable_file = WorkflowDraftVariableFile(
upload_file_id=upload_file.id,
@@ -1105,7 +1106,7 @@ class DraftVariableSaver:
length=original_length,
value_type=value_seg.value_type,
app_id=self._app_id,
tenant_id=self._user.current_tenant_id,
tenant_id=self._tenant_id,
user_id=self._user.id,
)
variable_file.id = str(uuidv7())
+90 -11
View File
@@ -13,6 +13,7 @@ from sqlalchemy import desc, select
from sqlalchemy.orm import Session, sessionmaker
from core.app.apps.message_generator import MessageGenerator
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity
from core.app.entities.task_entities import (
HumanInputRequiredResponse,
MessageReplaceStreamResponse,
@@ -84,9 +85,6 @@ def build_workflow_event_stream(
topic = MessageGenerator.get_response_topic(app_mode, workflow_run.id)
workflow_run_repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(session_maker)
node_execution_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository(session_maker)
message_context = (
_get_message_context(session_maker, workflow_run.id) if app_mode == AppMode.ADVANCED_CHAT else None
)
pause_entity: WorkflowPauseEntity | None = None
if workflow_run.status == WorkflowExecutionStatus.PAUSED:
@@ -97,6 +95,38 @@ def build_workflow_event_stream(
pause_entity = None
resumption_context = _load_resumption_context(pause_entity)
message_context: MessageContext | None = None
if app_mode == AppMode.ADVANCED_CHAT:
if workflow_run.status == WorkflowExecutionStatus.PAUSED:
if resumption_context is None:
raise AssertionError(
"WorkflowResumptionContext is required for advanced-chat snapshot replay, "
f"workflow_run_id={workflow_run.id}"
)
generate_entity = resumption_context.get_generate_entity()
if not isinstance(generate_entity, AdvancedChatAppGenerateEntity):
raise AssertionError(
"AdvancedChatAppGenerateEntity is required for advanced-chat snapshot replay, "
f"workflow_run_id={workflow_run.id}, generate_entity_type={type(generate_entity).__name__}"
)
if not generate_entity.conversation_id:
raise AssertionError(
f"conversation_id is required for advanced-chat snapshot replay, workflow_run_id={workflow_run.id}"
)
message_context = _get_message_context_by_conversation(
session_maker,
conversation_id=generate_entity.conversation_id,
workflow_run_id=workflow_run.id,
)
else:
# Compatibility fallback for non-suspended snapshot requests. This app-scoped lookup is not optimal;
# a dedicated index or stronger lookup key would be preferable.
message_context = _get_message_context_by_app(
session_maker,
app_id=app_id,
workflow_run_id=workflow_run.id,
)
node_snapshots = node_execution_repo.get_execution_snapshots_by_workflow_run(
tenant_id=tenant_id,
app_id=app_id,
@@ -175,19 +205,68 @@ def build_workflow_event_stream(
return _generate()
def _get_message_context(session_maker: sessionmaker[Session], workflow_run_id: str) -> MessageContext | None:
def _get_message_context_by_conversation(
session_maker: sessionmaker[Session],
*,
conversation_id: str,
workflow_run_id: str,
) -> MessageContext | None:
"""Look up a paused or suspended Advanced Chat snapshot message by conversation and workflow run.
Use this exact lookup after recovering ``conversation_id`` from persisted resumption context. Its predicates match
``message_workflow_run_id_idx``.
"""
with session_maker() as session:
stmt = select(Message).where(Message.workflow_run_id == workflow_run_id).order_by(desc(Message.created_at))
stmt = (
select(Message)
.where(
Message.conversation_id == conversation_id,
Message.workflow_run_id == workflow_run_id,
)
.order_by(desc(Message.created_at))
.limit(1)
)
message = session.scalar(stmt)
if message is None:
return None
created_at = int(message.created_at.timestamp()) if message.created_at else 0
return MessageContext(
conversation_id=message.conversation_id,
message_id=message.id,
created_at=created_at,
answer=message.answer,
return _to_message_context(message)
def _get_message_context_by_app(
session_maker: sessionmaker[Session],
*,
app_id: str,
workflow_run_id: str,
) -> MessageContext | None:
"""Look up a non-suspended or running Advanced Chat reconnect snapshot by app and workflow run.
This compatibility path applies only when no resumption context is expected. The app-scoped query is not optimal;
a dedicated index or stronger lookup key would be preferable.
"""
with session_maker() as session:
stmt = (
select(Message)
.where(
Message.app_id == app_id,
Message.workflow_run_id == workflow_run_id,
)
.order_by(desc(Message.created_at))
.limit(1)
)
message = session.scalar(stmt)
if message is None:
return None
return _to_message_context(message)
def _to_message_context(message: Message) -> MessageContext:
created_at = int(message.created_at.timestamp()) if message.created_at else 0
return MessageContext(
conversation_id=message.conversation_id,
message_id=message.id,
created_at=created_at,
answer=message.answer,
)
def _load_resumption_context(pause_entity: WorkflowPauseEntity | None) -> WorkflowResumptionContext | None:
+10 -42
View File
@@ -16,13 +16,11 @@ from collections.abc import Iterator
from typing import Any
from core.app.app_config.entities import ModelConfig
from core.llm_generator.llm_generator import LLMGenerator
from core.model_manager import ModelInstance, ModelManager
from core.workflow.generator import WorkflowGenerator
from core.workflow.generator.tool_catalogue import build_tool_catalogue, format_tool_catalogue, installed_tool_keys
from core.workflow.generator.types import (
WorkflowGenerateResultDict,
WorkflowGenerationMode,
WorkflowGenerationModeRequest,
)
from graphon.model_runtime.entities.model_entities import ModelType
@@ -52,11 +50,9 @@ class WorkflowGeneratorService:
"""
Resolve a model instance for the tenant and run the generator.
``mode`` accepts the ``"auto"`` sentinel when set, the instruction is
classified into a concrete ``workflow`` / ``advanced-chat`` mode (one
tiny LLM call) before planning so the rest of the pipeline runs against
a concrete mode. The resolved mode is echoed back under the result's
``mode`` key.
``mode`` accepts the ``"auto"`` sentinel the planner itself picks the
concrete ``workflow`` / ``advanced-chat`` mode (no extra LLM call) and
the resolution is echoed back under the result's ``mode`` key.
``current_graph`` is the existing draft graph for the cmd+k `/refine`
flow when present the generator refines it instead of creating a new
@@ -66,9 +62,6 @@ class WorkflowGeneratorService:
controller can map them to existing HTTP error envelopes (same
envelope as ``/rule-generate``).
"""
resolved_mode = cls._resolve_mode(
tenant_id=tenant_id, mode=mode, instruction=instruction, model_config=model_config
)
model_instance, model_parameters, tool_catalogue_text, installed_tools = cls._resolve_generation_context(
tenant_id=tenant_id, model_config=model_config
)
@@ -79,7 +72,7 @@ class WorkflowGeneratorService:
provider=model_config.provider,
model_name=model_config.name,
model_mode=model_config.mode.value,
mode=resolved_mode,
mode=mode,
instruction=instruction,
ideal_output=ideal_output,
tool_catalogue_text=tool_catalogue_text,
@@ -101,16 +94,13 @@ class WorkflowGeneratorService:
"""
Streaming sibling of ``generate_workflow_graph``.
Resolves the same model instance / tool catalogue / concrete mode, then
delegates to ``WorkflowGenerator.generate_workflow_graph_stream`` and
yields its ``(event_name, payload)`` tuples through to the controller's
SSE writer. Provider-init / invoke errors raised while resolving the
model instance propagate to the caller (the controller emits them as a
Resolves the same model instance / tool catalogue, then delegates to
``WorkflowGenerator.generate_workflow_graph_stream`` and yields its
``(event_name, payload)`` tuples through to the controller's SSE
writer. Provider-init / invoke errors raised while resolving the model
instance propagate to the caller (the controller emits them as a
single ``result`` SSE event).
"""
resolved_mode = cls._resolve_mode(
tenant_id=tenant_id, mode=mode, instruction=instruction, model_config=model_config
)
model_instance, model_parameters, tool_catalogue_text, installed_tools = cls._resolve_generation_context(
tenant_id=tenant_id, model_config=model_config
)
@@ -121,7 +111,7 @@ class WorkflowGeneratorService:
provider=model_config.provider,
model_name=model_config.name,
model_mode=model_config.mode.value,
mode=resolved_mode,
mode=mode,
instruction=instruction,
ideal_output=ideal_output,
tool_catalogue_text=tool_catalogue_text,
@@ -129,28 +119,6 @@ class WorkflowGeneratorService:
current_graph=current_graph,
)
@classmethod
def _resolve_mode(
cls,
*,
tenant_id: str,
mode: WorkflowGenerationModeRequest,
instruction: str,
model_config: ModelConfig,
) -> WorkflowGenerationMode:
"""Resolve the request mode into a concrete generation mode.
``"auto"`` triggers a one-word LLM classification using the model the
user already picked; everything else passes through unchanged. The
classifier never raises (defaults to ``advanced-chat``), so ``auto``
never blocks generation.
"""
if mode == "auto":
return LLMGenerator.classify_workflow_mode(
tenant_id=tenant_id, instruction=instruction, model_config=model_config
)
return mode
@classmethod
def _resolve_generation_context(
cls,
+11 -5
View File
@@ -322,6 +322,7 @@ class WorkflowService:
session: Session,
commit: bool = True,
sync_agent_bindings: bool = True,
graph_only: bool = False,
) -> Workflow:
"""
Sync draft workflow.
@@ -338,8 +339,10 @@ class WorkflowService:
if workflow and workflow.unique_hash != unique_hash:
raise WorkflowHashNotEqualError()
# validate features structure
self.validate_features_structure(app_model=app_model, features=features)
# Collaboration persists features and variables through dedicated endpoints. A graph save
# must not overwrite those newer database values with another collaborator's stale cache.
if not graph_only or not workflow:
self.validate_features_structure(app_model=app_model, features=features)
# validate graph structure
self.validate_graph_structure(graph=graph)
@@ -361,11 +364,12 @@ class WorkflowService:
# update draft workflow if found
else:
workflow.graph = json.dumps(graph)
workflow.features = json.dumps(features)
workflow.updated_by = account.id
workflow.updated_at = naive_utc_now()
workflow.environment_variables = environment_variables
workflow.conversation_variables = conversation_variables
if not graph_only:
workflow.features = json.dumps(features)
workflow.environment_variables = environment_variables
workflow.conversation_variables = conversation_variables
from services.agent.workflow_publish_service import WorkflowAgentPublishService
@@ -1056,6 +1060,7 @@ class WorkflowService:
with sessionmaker(bind=db.engine).begin() as session:
draft_var_saver = DraftVariableSaver(
session=session,
tenant_id=app_model.tenant_id,
app_id=app_model.id,
node_id=workflow_node_execution.node_id,
node_type=workflow_node_execution.node_type,
@@ -1206,6 +1211,7 @@ class WorkflowService:
with sessionmaker(bind=db.engine).begin() as session:
draft_var_saver = DraftVariableSaver(
session=session,
tenant_id=app_model.tenant_id,
app_id=app_model.id,
node_id=node_id,
node_type=BuiltinNodeTypes.HUMAN_INPUT,
@@ -24,6 +24,9 @@ def _create_agent_backend_client():
base_url=dify_config.AGENT_BACKEND_BASE_URL,
use_fake=dify_config.AGENT_BACKEND_USE_FAKE,
fake_scenario=dify_config.AGENT_BACKEND_FAKE_SCENARIO,
stream_read_timeout_seconds=dify_config.AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS,
stream_max_reconnects=dify_config.AGENT_BACKEND_STREAM_MAX_RECONNECTS,
stream_run_timeout_seconds=dify_config.AGENT_BACKEND_RUN_TIMEOUT_SECONDS,
)
+99
View File
@@ -0,0 +1,99 @@
"""Install configured marketplace plugins for a newly created tenant."""
import logging
from celery import shared_task
from configs import dify_config
from core.helper import marketplace
from core.plugin.entities.plugin_daemon import PluginInstallTaskStatus
from core.plugin.plugin_service import PluginService
from services.model_provider_service import ModelProviderService
logger = logging.getLogger(__name__)
@shared_task(queue="plugin", bind=True, max_retries=60, default_retry_delay=5)
def configure_default_models_task(self, tenant_id: str, plugin_install_task_id: str | None) -> None:
"""Set explicitly configured default models after default plugins finish installing."""
if not dify_config.NEW_USER_DEFAULT_MODELS:
return
plugin_install_failed = False
if plugin_install_task_id:
try:
install_task = PluginService.fetch_install_task(tenant_id, plugin_install_task_id)
except Exception as exc:
logger.warning(
"Failed to fetch default plugin installation task for tenant %s; retrying",
tenant_id,
)
raise self.retry(exc=exc)
if install_task.status in (PluginInstallTaskStatus.Pending, PluginInstallTaskStatus.Running):
raise self.retry()
plugin_install_failed = install_task.status == PluginInstallTaskStatus.Failed
if plugin_install_failed:
failed_plugin_ids = [
plugin.plugin_id for plugin in install_task.plugins if plugin.status == PluginInstallTaskStatus.Failed
]
logger.error(
"Default plugin installation failed for tenant %s: %s",
tenant_id,
", ".join(failed_plugin_ids),
)
model_provider_service = ModelProviderService()
failed_model_types: list[str] = []
for model_type, provider, model in dify_config.NEW_USER_DEFAULT_MODEL_LIST:
try:
model_provider_service.update_default_model_of_model_type(
tenant_id=tenant_id,
model_type=model_type,
provider=provider,
model=model,
)
except Exception:
failed_model_types.append(model_type)
logger.exception(
"Failed to configure default model for tenant %s: model_type=%s provider=%s model=%s",
tenant_id,
model_type,
provider,
model,
)
if plugin_install_failed or failed_model_types:
raise RuntimeError(
f"Failed to initialize defaults for tenant {tenant_id}; "
f"model types: {', '.join(failed_model_types) or 'none'}"
)
@shared_task(queue="plugin")
def install_default_plugins_task(tenant_id: str, plugin_ids: list[str]) -> None:
"""Install the latest marketplace versions of the configured plugins."""
if not plugin_ids:
return
try:
manifests = {manifest.plugin_id: manifest for manifest in marketplace.batch_fetch_plugin_manifests(plugin_ids)}
plugin_identifiers = [
manifests[plugin_id].latest_package_identifier for plugin_id in plugin_ids if plugin_id in manifests
]
missing_plugin_ids = [plugin_id for plugin_id in plugin_ids if plugin_id not in manifests]
if missing_plugin_ids:
logger.warning("Default plugins not found in marketplace: %s", ", ".join(missing_plugin_ids))
if not plugin_identifiers:
return
response = PluginService.install_from_marketplace_pkg(tenant_id, plugin_identifiers)
if dify_config.NEW_USER_DEFAULT_MODELS:
configure_default_models_task.delay(
tenant_id,
None if response.all_installed else response.task_id,
)
except Exception:
logger.exception("Failed to install default plugins for tenant %s", tenant_id)
raise
+1
View File
@@ -1,5 +1,6 @@
FLASK_APP=app.py
FLASK_DEBUG=0
ENABLE_COLLABORATION_MODE=false
SECRET_KEY='uhySf6a3aZuvRNfAlcr47paOw9TRYBY6j8ZHXpVw1yx5RP27Yj3w2uvI'
CONSOLE_API_URL=http://127.0.0.1:5001
@@ -513,11 +513,94 @@ class TestArchiveRunIdempotency:
storage = MagicMock()
storage.object_exists.return_value = True
result = archiver._archive_bundle(MagicMock(), storage, [run])
with patch.object(archiver, "_sync_existing_bundle_index") as sync_existing_bundle_index:
result = archiver._archive_bundle(MagicMock(), storage, [run])
assert result.success is True
assert result.skipped is True
assert result.error == "bundle already archived"
sync_existing_bundle_index.assert_called_once()
def test_existing_bundle_catalog_publication_failure_is_not_success(self):
archiver = WorkflowRunArchiver(days=90)
run = _run()
session = MagicMock()
storage = MagicMock()
storage.object_exists.return_value = True
with patch.object(archiver, "_sync_existing_bundle_index", side_effect=RuntimeError("catalog unavailable")):
result = archiver._archive_bundle(session, storage, [run])
assert result.success is False
assert result.error == "catalog unavailable"
session.rollback.assert_called_once()
def test_retry_repairs_index_after_catalog_commit_then_index_write_failure(self):
archiver = WorkflowRunArchiver(days=90)
run = _run()
identity = archiver._build_bundle_identity([run])
index_key = archiver._get_index_object_key(identity)
manifest_key = archiver._get_manifest_object_key(identity)
storage = FakeArchiveStorage()
original_put_object = storage.put_object
index_write_count = 0
def put_object(key: str, data: bytes) -> str:
nonlocal index_write_count
if key == index_key:
index_write_count += 1
if index_write_count == 2:
raise RuntimeError("index write failed")
return original_put_object(key, data)
storage.put_object = MagicMock(side_effect=put_object)
first_session = MagicMock()
first_session.scalar.return_value = None
table_data = {"workflow_runs": [{"id": run.id, "tenant_id": run.tenant_id}]}
with (
patch.object(archiver, "_lock_runs_for_archive", return_value=[run]),
patch.object(archiver, "_extract_bundle_data", return_value=table_data),
):
first_result = archiver._archive_bundle(first_session, storage, [run])
assert first_result.success is False
assert first_result.error == "index write failed"
assert manifest_key in storage.objects
assert json.loads(storage.objects[index_key])["run_ids"] == []
storage.list_objects = MagicMock(wraps=storage.list_objects)
retry_session = MagicMock()
retry_session.scalar.return_value = None
retry_result = archiver._archive_bundle(retry_session, storage, [run])
assert retry_result.success is True
assert retry_result.skipped is True
assert json.loads(storage.objects[index_key])["manifest_keys"] == [manifest_key]
assert json.loads(storage.objects[index_key])["run_ids"] == [run.id]
storage.list_objects.assert_not_called()
def test_existing_manifest_with_missing_index_fails_without_partial_rebuild(self):
archiver = WorkflowRunArchiver(days=90)
run = _run()
identity = archiver._build_bundle_identity([run])
_, _, manifest_data = archiver._build_archive_payload(
identity,
[run],
{"workflow_runs": [{"id": run.id, "tenant_id": run.tenant_id}]},
)
manifest_key = archiver._get_manifest_object_key(identity)
index_key = archiver._get_index_object_key(identity)
storage = FakeArchiveStorage({manifest_key: manifest_data})
storage.list_objects = MagicMock(wraps=storage.list_objects)
result = archiver._archive_bundle(MagicMock(), storage, [run])
assert result.success is False
assert "archive shard index missing" in (result.error or "")
assert index_key not in storage.objects
storage.list_objects.assert_not_called()
def test_successful_bundle_persists_archive_index(self):
archiver = WorkflowRunArchiver(days=90)
@@ -550,6 +633,28 @@ class TestArchiveRunIdempotency:
assert archived_bundle.row_count == 2
session.commit.assert_called_once()
def test_new_bundle_catalog_commit_failure_is_not_success(self):
archiver = WorkflowRunArchiver(days=90)
run = _run(str(uuid.uuid4()))
run.tenant_id = str(uuid.uuid4())
session = MagicMock()
session.scalar.return_value = None
session.commit.side_effect = RuntimeError("catalog commit failed")
storage = MagicMock()
storage.object_exists.return_value = False
storage.list_objects.return_value = []
table_data = {"workflow_runs": [{"id": run.id, "tenant_id": run.tenant_id}]}
with (
patch.object(archiver, "_lock_runs_for_archive", return_value=[run]),
patch.object(archiver, "_extract_bundle_data", return_value=table_data),
):
result = archiver._archive_bundle(session, storage, [run])
assert result.success is False
assert result.error == "catalog commit failed"
session.rollback.assert_called_once()
def test_index_skips_all_already_archived_runs(self):
archiver = WorkflowRunArchiver(days=90)
run = MagicMock()
@@ -311,6 +311,7 @@ class TestDraftVariableLoader(unittest.TestCase):
# Use DraftVariableSaver to create offloaded variable (this mimics production)
saver = DraftVariableSaver(
session=session,
tenant_id=self._test_tenant_id,
app_id=self._test_app_id,
node_id="test_offload_node",
node_type=BuiltinNodeTypes.LLM, # Use a real node type
@@ -1,179 +0,0 @@
"""Testcontainers integration tests for controllers.console.app.app_import endpoints."""
from __future__ import annotations
from inspect import unwrap
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from flask import Flask
from controllers.console.app import app_import as app_import_module
from services.app_dsl_service import ImportStatus
class _Result:
def __init__(self, status: ImportStatus, app_id: str | None = "app-1"):
self.status = status
self.app_id = app_id
def model_dump(self, mode: str = "json"):
return {"status": self.status, "app_id": self.app_id}
def _install_features(monkeypatch: pytest.MonkeyPatch, enabled: bool) -> None:
features = SimpleNamespace(webapp_auth=SimpleNamespace(enabled=enabled))
monkeypatch.setattr(app_import_module.FeatureService, "get_system_features", lambda: features)
class TestAppImportApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask):
return flask_app_with_containers
def test_import_post_returns_failed_status(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportApi()
method = unwrap(api.post)
_install_features(monkeypatch, enabled=False)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.FAILED, app_id=None),
)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method(api, SimpleNamespace(id="u1"))
assert status == 400
assert response["status"] == ImportStatus.FAILED
def test_import_post_returns_pending_status(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportApi()
method = unwrap(api.post)
_install_features(monkeypatch, enabled=False)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.PENDING),
)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method(api, SimpleNamespace(id="u1"))
assert status == 202
assert response["status"] == ImportStatus.PENDING
def test_import_post_updates_webapp_auth_when_enabled(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportApi()
method = unwrap(api.post)
_install_features(monkeypatch, enabled=True)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.COMPLETED, app_id="app-123"),
)
update_access = MagicMock()
monkeypatch.setattr(app_import_module.EnterpriseService.WebAppAuth, "update_app_access_mode", update_access)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method(api, SimpleNamespace(id="u1"))
update_access.assert_called_once_with("app-123", "private")
assert status == 200
assert response["status"] == ImportStatus.COMPLETED
def test_import_post_commits_session_on_success(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportApi()
method = unwrap(api.post)
_install_features(monkeypatch, enabled=False)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.COMPLETED, app_id="app-123"),
)
fake_session = MagicMock()
fake_session.__enter__.return_value = fake_session
fake_session.__exit__.return_value = None
monkeypatch.setattr(app_import_module, "Session", lambda *_args, **_kwargs: fake_session)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method(api, SimpleNamespace(id="u1"))
fake_session.commit.assert_called_once_with()
fake_session.rollback.assert_not_called()
assert status == 200
assert response["status"] == ImportStatus.COMPLETED
def test_import_post_rolls_back_session_on_failure(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportApi()
method = unwrap(api.post)
_install_features(monkeypatch, enabled=False)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.FAILED, app_id=None),
)
fake_session = MagicMock()
fake_session.__enter__.return_value = fake_session
fake_session.__exit__.return_value = None
monkeypatch.setattr(app_import_module, "Session", lambda *_args, **_kwargs: fake_session)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method(api, SimpleNamespace(id="u1"))
fake_session.rollback.assert_called_once_with()
fake_session.commit.assert_not_called()
assert status == 400
assert response["status"] == ImportStatus.FAILED
class TestAppImportConfirmApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask):
return flask_app_with_containers
def test_import_confirm_returns_failed_status(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportConfirmApi()
method = unwrap(api.post)
monkeypatch.setattr(
app_import_module.AppDslService,
"confirm_import",
lambda *_args, **_kwargs: _Result(ImportStatus.FAILED),
)
with app.test_request_context("/console/api/apps/imports/import-1/confirm", method="POST"):
response, status = method(api, SimpleNamespace(id="u1"), import_id="import-1")
assert status == 400
assert response["status"] == ImportStatus.FAILED
class TestAppImportCheckDependenciesApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask):
return flask_app_with_containers
def test_import_check_dependencies_returns_result(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportCheckDependenciesApi()
method = unwrap(api.get)
monkeypatch.setattr(
app_import_module.AppDslService,
"check_dependencies",
lambda *_args, **_kwargs: SimpleNamespace(model_dump=lambda mode="json": {"leaked_dependencies": []}),
)
with app.test_request_context("/console/api/apps/imports/app-1/check-dependencies", method="GET"):
response, status = method(api, app_model=SimpleNamespace(id="app-1"))
assert status == 200
assert response["leaked_dependencies"] == []
@@ -1,4 +1,4 @@
from collections.abc import Iterator
from collections.abc import Callable, Iterator
from typing import override
import pytest
@@ -47,6 +47,8 @@ def _request():
class _SuccessfulClient:
stream_options: tuple[int | None, float | None, Callable[[], bool] | None] | None = None
def create_run_sync(self, request: CreateRunRequest) -> CreateRunResponse:
assert isinstance(request, CreateRunRequest)
return CreateRunResponse(run_id="run-1", status="running")
@@ -55,8 +57,17 @@ class _SuccessfulClient:
del request
return CancelRunResponse(run_id=run_id, status="cancelled")
def stream_events_sync(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]:
def stream_events_sync(
self,
run_id: str,
*,
after: str | None = None,
max_reconnects: int | None = None,
timeout_seconds: float | None = None,
should_stop: Callable[[], bool] | None = None,
) -> Iterator[RunEvent]:
del after
self.stream_options = (max_reconnects, timeout_seconds, should_stop)
yield RunStartedEvent(id="1-0", run_id=run_id)
def wait_run_sync(self, run_id: str, *, timeout_seconds: float | None = None) -> RunStatusResponse:
@@ -72,17 +83,22 @@ class _SuccessfulClient:
def test_dify_agent_backend_run_client_delegates_sync_methods():
client = DifyAgentBackendRunClient(_SuccessfulClient())
wrapped = _SuccessfulClient()
client = DifyAgentBackendRunClient(wrapped, stream_max_reconnects=2, stream_timeout_seconds=45)
def should_stop() -> bool:
return False
created = client.create_run(_request())
cancelled = client.cancel_run(created.run_id)
events = list(client.stream_events(created.run_id))
events = list(client.stream_events(created.run_id, should_stop=should_stop))
status = client.wait_run(created.run_id)
assert created.run_id == "run-1"
assert cancelled.status == "cancelled"
assert events[0].type == "run_started"
assert status.status == "succeeded"
assert wrapped.stream_options == (2, 45, should_stop)
def test_dify_agent_backend_run_client_maps_validation_error():
@@ -125,7 +141,16 @@ def test_dify_agent_backend_run_client_maps_timeout_error():
def test_dify_agent_backend_run_client_maps_stream_error():
class StreamClient(_SuccessfulClient):
@override
def stream_events_sync(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]:
def stream_events_sync(
self,
run_id: str,
*,
after: str | None = None,
max_reconnects: int | None = None,
timeout_seconds: float | None = None,
should_stop: Callable[[], bool] | None = None,
) -> Iterator[RunEvent]:
del run_id, after, max_reconnects, timeout_seconds, should_stop
raise DifyAgentStreamError("bad stream")
yield
@@ -1,3 +1,5 @@
from decimal import Decimal
import pytest
from agenton.compositor import CompositorSessionSnapshot
from dify_agent.protocol import (
@@ -83,7 +85,19 @@ def test_event_adapter_maps_run_succeeded_to_final_output():
data=RunSucceededEventData(
output={"summary": "done"},
session_snapshot=snapshot,
usage=AgentRunUsage(prompt_tokens=2, completion_tokens=3),
usage=AgentRunUsage(
prompt_tokens=2,
prompt_unit_price=Decimal(5),
prompt_price_unit=Decimal("0.000001"),
prompt_price=Decimal("0.000010"),
completion_tokens=3,
completion_unit_price=Decimal(30),
completion_price_unit=Decimal("0.000001"),
completion_price=Decimal("0.000090"),
total_price=Decimal("0.000100"),
currency="USD",
latency=0.4,
),
),
)
)
@@ -94,7 +108,22 @@ def test_event_adapter_maps_run_succeeded_to_final_output():
source_event_id="3-0",
output={"summary": "done"},
session_snapshot=snapshot,
usage={"prompt_tokens": 2, "completion_tokens": 3, "total_tokens": 5},
usage={
"prompt_tokens": 2,
"prompt_unit_price": "5",
"prompt_price_unit": "0.000001",
"prompt_price": "0.000010",
"completion_tokens": 3,
"completion_unit_price": "30",
"completion_price_unit": "0.000001",
"completion_price": "0.000090",
"total_tokens": 5,
"total_price": "0.000100",
"currency": "USD",
"latency": 0.4,
"time_to_first_token": None,
"time_to_generate": None,
},
)
]
@@ -3,9 +3,19 @@ from unittest.mock import MagicMock
import click
import pytest
from click.testing import CliRunner
from sqlalchemy.exc import OperationalError
from commands import retention
from services.retention.workflow_run import bundle_archive_maintenance
from services.retention.workflow_run.bundle_archive_maintenance import (
BundleOperationResult,
BundleOperationSummary,
)
_CURSOR_0 = "00000000-0000-0000-0000-000000000000"
_CURSOR_1 = "00000000-0000-0000-0000-000000000001"
_CURSOR_2 = "00000000-0000-0000-0000-000000000002"
def _db_disconnect_error() -> OperationalError:
@@ -24,6 +34,58 @@ def _session_context(session):
return context
def _delete_summary(
*,
processed: int,
succeeded: int = 0,
failed: int = 0,
next_catalog_id: str | None = None,
preview_next_catalog_id: str | None = None,
results: list[BundleOperationResult] | None = None,
) -> BundleOperationSummary:
return BundleOperationSummary(
operation="delete",
bundles_processed=processed,
bundles_succeeded=succeeded,
bundles_failed=failed,
next_catalog_id=next_catalog_id,
preview_next_catalog_id=preview_next_catalog_id,
results=results or [],
)
def _patch_bundle_deleter(
monkeypatch: pytest.MonkeyPatch,
summaries: list[BundleOperationSummary],
) -> MagicMock:
deleter = MagicMock()
deleter.delete_batch.side_effect = summaries
monkeypatch.setattr(
bundle_archive_maintenance,
"WorkflowRunBundleArchiveMaintenance",
MagicMock(return_value=deleter),
)
return deleter
@pytest.mark.parametrize(
"command",
[retention.restore_workflow_runs, retention.delete_archived_workflow_runs],
)
def test_v2_archive_maintenance_rejects_explicitly_empty_tenant_ids(command):
result = CliRunner().invoke(
command,
["--tenant-ids", "", "--target-month", "2025-03"],
)
assert result.exit_code == 2
assert "tenant-ids must not be empty" in result.output
def test_archive_tenant_id_parser_keeps_omitted_scope_unset():
assert retention._parse_comma_separated_ids(None, param_name="tenant-ids") is None
def test_resolve_archive_tenant_ids_from_plan_uses_explicit_sessions(monkeypatch):
end_before = datetime.datetime(2025, 4, 1, tzinfo=datetime.UTC)
sessions = [MagicMock(name="session-a"), MagicMock(name="session-b")]
@@ -137,3 +199,265 @@ def test_archive_workflow_runs_raises_click_exception_when_tenant_plan_fails(mon
dry_run=True,
delete_after_archive=False,
)
def test_delete_archived_workflow_runs_keeps_single_page_behavior_without_all_pages(monkeypatch):
deleter = _patch_bundle_deleter(
monkeypatch,
[_delete_summary(processed=2, succeeded=2, next_catalog_id=_CURSOR_1)],
)
result = CliRunner().invoke(
retention.delete_archived_workflow_runs,
["--target-month", "2025-03", "--limit", "2"],
)
assert result.exit_code == 0
deleter.delete_batch.assert_called_once()
assert deleter.delete_batch.call_args.kwargs["target_year"] == 2025
assert deleter.delete_batch.call_args.kwargs["target_month"] == 3
assert deleter.delete_batch.call_args.kwargs["after_catalog_id"] is None
assert deleter.delete_batch.call_args.kwargs["limit"] == 2
def test_delete_archived_workflow_runs_all_pages_continues_until_empty_page(monkeypatch):
deleter = _patch_bundle_deleter(
monkeypatch,
[
_delete_summary(processed=2, succeeded=2, next_catalog_id=_CURSOR_1),
_delete_summary(processed=1, succeeded=1, next_catalog_id=_CURSOR_2),
_delete_summary(processed=0),
],
)
result = CliRunner().invoke(
retention.delete_archived_workflow_runs,
["--target-month", "2025-03", "--all-pages", "--limit", "2"],
)
assert result.exit_code == 0
assert [call.kwargs["after_catalog_id"] for call in deleter.delete_batch.call_args_list] == [
None,
_CURSOR_1,
_CURSOR_2,
]
def test_delete_archived_workflow_runs_all_pages_fetches_empty_page_after_exact_full_page(monkeypatch):
deleter = _patch_bundle_deleter(
monkeypatch,
[
_delete_summary(processed=2, succeeded=2, next_catalog_id=_CURSOR_1),
_delete_summary(processed=0),
],
)
result = CliRunner().invoke(
retention.delete_archived_workflow_runs,
["--target-month", "2025-03", "--all-pages", "--limit", "2"],
)
assert result.exit_code == 0
assert deleter.delete_batch.call_count == 2
assert deleter.delete_batch.call_args_list[1].kwargs["after_catalog_id"] == _CURSOR_1
def test_delete_archived_workflow_runs_all_pages_stops_at_first_failed_page(monkeypatch):
failed_result = BundleOperationResult(
catalog_id=_CURSOR_2,
bundle_id="bundle-failed",
tenant_id="tenant-1",
object_prefix="workflow-runs/v2/tenant-1/2025/03/00-of-16/bundle-failed",
error="archive checksum mismatch",
)
deleter = _patch_bundle_deleter(
monkeypatch,
[
_delete_summary(processed=1, succeeded=1, next_catalog_id=_CURSOR_1),
_delete_summary(processed=1, failed=1, results=[failed_result]),
_delete_summary(processed=0),
],
)
result = CliRunner().invoke(
retention.delete_archived_workflow_runs,
["--target-month", "2025-03", "--all-pages"],
)
assert result.exit_code == 1
assert deleter.delete_batch.call_count == 2
assert "target_month=2025-03" in result.output
assert f"failed_catalog_id={_CURSOR_2}" in result.output
assert f"resume_after_catalog_id={_CURSOR_1}" in result.output
def test_delete_archived_workflow_runs_all_pages_fails_when_cursor_does_not_advance(monkeypatch):
deleter = _patch_bundle_deleter(
monkeypatch,
[_delete_summary(processed=1, succeeded=1, next_catalog_id=None)],
)
result = CliRunner().invoke(
retention.delete_archived_workflow_runs,
["--target-month", "2025-03", "--all-pages"],
)
assert result.exit_code == 1
deleter.delete_batch.assert_called_once()
assert "cursor did not advance" in result.output.lower()
def test_delete_archived_workflow_runs_all_pages_uses_preview_cursor_for_dry_run(monkeypatch):
deleter = _patch_bundle_deleter(
monkeypatch,
[
_delete_summary(processed=1, succeeded=1, preview_next_catalog_id=_CURSOR_1),
_delete_summary(processed=0),
],
)
result = CliRunner().invoke(
retention.delete_archived_workflow_runs,
["--target-month", "2025-03", "--all-pages", "--dry-run"],
)
assert result.exit_code == 0
assert [call.kwargs["after_catalog_id"] for call in deleter.delete_batch.call_args_list] == [
None,
_CURSOR_1,
]
def test_delete_archived_workflow_runs_dry_run_failure_separates_preview_and_destructive_cursors(monkeypatch):
failed_result = BundleOperationResult(
catalog_id=_CURSOR_2,
bundle_id="bundle-failed",
tenant_id="tenant-1",
object_prefix="workflow-runs/v2/tenant-1/2025/03/00-of-16/bundle-failed",
error="archive checksum mismatch",
)
deleter = _patch_bundle_deleter(
monkeypatch,
[
_delete_summary(
processed=1,
succeeded=1,
preview_next_catalog_id=_CURSOR_1,
),
_delete_summary(
processed=1,
failed=1,
results=[failed_result],
),
],
)
result = CliRunner().invoke(
retention.delete_archived_workflow_runs,
[
"--target-month",
"2025-03",
"--after-catalog-id",
_CURSOR_0,
"--all-pages",
"--dry-run",
],
)
assert result.exit_code == 1
assert deleter.delete_batch.call_count == 2
assert f"failed_catalog_id={_CURSOR_2}" in result.output
assert f"preview_after_catalog_id={_CURSOR_1}" in result.output
assert f"destructive_retry_after_catalog_id={_CURSOR_0}" in result.output
def test_delete_archived_workflow_runs_all_pages_starts_after_explicit_cursor(monkeypatch):
deleter = _patch_bundle_deleter(monkeypatch, [_delete_summary(processed=0)])
result = CliRunner().invoke(
retention.delete_archived_workflow_runs,
[
"--target-month",
"2025-03",
"--after-catalog-id",
_CURSOR_0,
"--all-pages",
],
)
assert result.exit_code == 0
deleter.delete_batch.assert_called_once()
assert deleter.delete_batch.call_args.kwargs["after_catalog_id"] == _CURSOR_0
@pytest.mark.parametrize(
"shard_args",
[
["--run-shard-index", "0"],
["--run-shard-total", "16"],
["--run-shard-index", "16", "--run-shard-total", "16"],
["--run-shard-index", "-1", "--run-shard-total", "16"],
["--run-shard-index", "0", "--run-shard-total", "0"],
["--run-shard-index", "0", "--run-shard-total", "17"],
],
)
def test_delete_archived_workflow_runs_rejects_invalid_run_shard_options(monkeypatch, shard_args):
deleter = _patch_bundle_deleter(monkeypatch, [_delete_summary(processed=0)])
result = CliRunner().invoke(
retention.delete_archived_workflow_runs,
["--target-month", "2025-03", *shard_args],
)
assert result.exit_code == 2
deleter.delete_batch.assert_not_called()
def test_delete_archived_workflow_runs_passes_formatted_run_shard_to_service(monkeypatch):
deleter = _patch_bundle_deleter(monkeypatch, [_delete_summary(processed=0)])
result = CliRunner().invoke(
retention.delete_archived_workflow_runs,
[
"--target-month",
"2025-03",
"--tenant-ids",
"tenant-1",
"--run-shard-index",
"3",
"--run-shard-total",
"16",
],
)
assert result.exit_code == 0
deleter.validate_catalog_shards.assert_called_once_with(
target_year=2025,
target_month=3,
shard_total=16,
tenant_ids=["tenant-1"],
)
deleter.delete_batch.assert_called_once()
assert deleter.delete_batch.call_args.kwargs["shard"] == "03-of-16"
def test_delete_archived_workflow_runs_rejects_mixed_catalog_shards_before_delete(monkeypatch):
deleter = _patch_bundle_deleter(monkeypatch, [_delete_summary(processed=0)])
deleter.validate_catalog_shards.side_effect = ValueError("unexpected shards: 00-of-01")
result = CliRunner().invoke(
retention.delete_archived_workflow_runs,
[
"--target-month",
"2025-03",
"--run-shard-index",
"3",
"--run-shard-total",
"16",
],
)
assert result.exit_code == 1
assert "shard preflight failed" in result.output.lower()
assert "00-of-01" in result.output
deleter.delete_batch.assert_not_called()
@@ -8,8 +8,13 @@ from yarl import URL
from configs.app_config import DifyConfig
def _clear_environment(monkeypatch: pytest.MonkeyPatch) -> None:
for name in tuple(os.environ):
monkeypatch.delenv(name)
def _set_basic_config_env(monkeypatch: pytest.MonkeyPatch) -> None:
os.environ.clear()
_clear_environment(monkeypatch)
monkeypatch.setenv("CONSOLE_API_URL", "https://example.com")
monkeypatch.setenv("CONSOLE_WEB_URL", "https://example.com")
monkeypatch.setenv("DB_TYPE", "postgresql")
@@ -51,7 +56,7 @@ def test_dify_config_preserves_explicit_secret_key(
def test_dify_config(monkeypatch: pytest.MonkeyPatch):
# clear system environment variables
os.environ.clear()
_clear_environment(monkeypatch)
# Set environment variables using monkeypatch
monkeypatch.setenv("CONSOLE_API_URL", "https://example.com")
@@ -89,10 +94,76 @@ def test_dify_config(monkeypatch: pytest.MonkeyPatch):
assert Version(config.project.version) >= Version("1.0.0")
def test_new_user_default_plugin_ids_are_parsed_from_env(monkeypatch: pytest.MonkeyPatch) -> None:
_set_basic_config_env(monkeypatch)
monkeypatch.setenv(
"NEW_USER_DEFAULT_PLUGIN_IDS",
"langgenius/openai, langgenius/gemini",
)
config = DifyConfig(_env_file=None)
assert config.NEW_USER_DEFAULT_PLUGIN_ID_LIST == [
"langgenius/openai",
"langgenius/gemini",
]
def test_plugin_remote_install_port_rejects_host_port_spec(monkeypatch: pytest.MonkeyPatch) -> None:
"""A 'host:port' compose publish spec must produce an actionable error, not an opaque int_parsing traceback."""
_set_basic_config_env(monkeypatch)
monkeypatch.setenv("PLUGIN_REMOTE_INSTALL_PORT", "127.0.0.1:5003")
with pytest.raises(ValueError, match="must be a bare port number"):
DifyConfig(_env_file=None)
def test_plugin_remote_install_port_accepts_bare_port(monkeypatch: pytest.MonkeyPatch) -> None:
_set_basic_config_env(monkeypatch)
monkeypatch.setenv("PLUGIN_REMOTE_INSTALL_PORT", "5003")
config = DifyConfig(_env_file=None)
assert config.PLUGIN_REMOTE_INSTALL_PORT == 5003
def test_new_user_default_models_are_parsed_from_env(monkeypatch: pytest.MonkeyPatch) -> None:
_set_basic_config_env(monkeypatch)
monkeypatch.setenv(
"NEW_USER_DEFAULT_MODELS",
(
"llm:langgenius/openai/openai:gpt-4o-mini, "
"text-embedding:langgenius/openai/openai:text-embedding-3-small, "
"rerank:langgenius/ollama/ollama:reranker:latest"
),
)
config = DifyConfig(_env_file=None)
assert config.NEW_USER_DEFAULT_MODEL_LIST == [
("llm", "langgenius/openai/openai", "gpt-4o-mini"),
("text-embedding", "langgenius/openai/openai", "text-embedding-3-small"),
("rerank", "langgenius/ollama/ollama", "reranker:latest"),
]
def test_new_user_default_models_reject_duplicate_model_types(monkeypatch: pytest.MonkeyPatch) -> None:
_set_basic_config_env(monkeypatch)
monkeypatch.setenv(
"NEW_USER_DEFAULT_MODELS",
"llm:langgenius/openai/openai:gpt-4o-mini,llm:langgenius/anthropic/anthropic:claude-sonnet-4",
)
config = DifyConfig(_env_file=None)
with pytest.raises(ValueError, match="duplicate model type: llm"):
_ = config.NEW_USER_DEFAULT_MODEL_LIST
def test_http_timeout_defaults(monkeypatch: pytest.MonkeyPatch):
"""Test that HTTP timeout defaults are correctly set"""
# clear system environment variables
os.environ.clear()
_clear_environment(monkeypatch)
# Set minimal required env vars
monkeypatch.setenv("DB_TYPE", "postgresql")
@@ -112,7 +183,7 @@ def test_http_timeout_defaults(monkeypatch: pytest.MonkeyPatch):
def test_internal_files_url_falls_back_to_server_console_api_url(monkeypatch: pytest.MonkeyPatch):
os.environ.clear()
_clear_environment(monkeypatch)
monkeypatch.setenv("SERVER_CONSOLE_API_URL", "http://api:5001")
config = DifyConfig(_env_file=None)
@@ -121,7 +192,7 @@ def test_internal_files_url_falls_back_to_server_console_api_url(monkeypatch: py
def test_internal_files_url_prefers_explicit_value(monkeypatch: pytest.MonkeyPatch):
os.environ.clear()
_clear_environment(monkeypatch)
monkeypatch.setenv("INTERNAL_FILES_URL", "http://files-internal:5001")
monkeypatch.setenv("SERVER_CONSOLE_API_URL", "http://api:5001")
@@ -135,7 +206,7 @@ def test_internal_files_url_prefers_explicit_value(monkeypatch: pytest.MonkeyPat
def test_flask_configs(monkeypatch: pytest.MonkeyPatch):
flask_app = Flask("app")
# clear system environment variables
os.environ.clear()
_clear_environment(monkeypatch)
# Set environment variables using monkeypatch
monkeypatch.setenv("CONSOLE_API_URL", "https://example.com")
@@ -242,7 +313,7 @@ def test_db_session_timezone_override_can_disable_app_level_timezone_injection(m
def test_pubsub_redis_url_default(monkeypatch: pytest.MonkeyPatch):
os.environ.clear()
_clear_environment(monkeypatch)
monkeypatch.setenv("CONSOLE_API_URL", "https://example.com")
monkeypatch.setenv("CONSOLE_WEB_URL", "https://example.com")
@@ -265,7 +336,7 @@ def test_pubsub_redis_url_default(monkeypatch: pytest.MonkeyPatch):
def test_pubsub_redis_url_override(monkeypatch: pytest.MonkeyPatch):
os.environ.clear()
_clear_environment(monkeypatch)
monkeypatch.setenv("CONSOLE_API_URL", "https://example.com")
monkeypatch.setenv("CONSOLE_WEB_URL", "https://example.com")
@@ -282,7 +353,7 @@ def test_pubsub_redis_url_override(monkeypatch: pytest.MonkeyPatch):
def test_pubsub_redis_url_required_when_default_unavailable(monkeypatch: pytest.MonkeyPatch):
os.environ.clear()
_clear_environment(monkeypatch)
monkeypatch.setenv("CONSOLE_API_URL", "https://example.com")
monkeypatch.setenv("CONSOLE_WEB_URL", "https://example.com")
@@ -298,7 +369,7 @@ def test_pubsub_redis_url_required_when_default_unavailable(monkeypatch: pytest.
def test_dify_config_exposes_redis_key_prefix_default(monkeypatch: pytest.MonkeyPatch):
os.environ.clear()
_clear_environment(monkeypatch)
monkeypatch.setenv("CONSOLE_API_URL", "https://example.com")
monkeypatch.setenv("CONSOLE_WEB_URL", "https://example.com")
@@ -315,7 +386,7 @@ def test_dify_config_exposes_redis_key_prefix_default(monkeypatch: pytest.Monkey
def test_dify_config_reads_redis_key_prefix_from_env(monkeypatch: pytest.MonkeyPatch):
os.environ.clear()
_clear_environment(monkeypatch)
monkeypatch.setenv("CONSOLE_API_URL", "https://example.com")
monkeypatch.setenv("CONSOLE_WEB_URL", "https://example.com")
@@ -367,7 +438,7 @@ def test_celery_broker_url_with_special_chars_password(
from kombu.utils.url import parse_url
# clear system environment variables
os.environ.clear()
_clear_environment(monkeypatch)
# Set up basic required environment variables (following existing pattern)
monkeypatch.setenv("CONSOLE_API_URL", "https://example.com")

Some files were not shown because too many files have changed in this diff Show More