Compare commits

...
Author SHA1 Message Date
hjlarry ecea749297 Revert "chore: improve lambda plugin runtime error display"
This reverts commit 6bc1508a6b.
2026-07-23 13:35:52 +08:00
hjlarry 31fc6c71f7 fix(api): prevent identity logging deadlock 2026-07-23 13:16:38 +08:00
hjlarry 08f08c61c8 test(api): re-enable identity context log filter 2026-07-23 13:16:32 +08:00
hjlarry 6bc1508a6b chore: improve lambda plugin runtime error display 2026-07-23 09:17:07 +08:00
hjlarry be697e3050 test(api): disable identity context log filter 2026-07-23 09:16:22 +08:00
林玮 (Jade Lin) 688a279a75 fix(api): use resource tenant for draft variable files (#39307)
(cherry picked from commit 28d17da3b1)
2026-07-21 10:24:49 +08:00
林玮 (Jade Lin) 245876c468 fix(trigger): return 429 instead of 500 when API quota is exceeded (#38664)
(cherry picked from commit 4043adacd7)
2026-07-20 15:46:46 +08:00
林玮 (Jade Lin) 82c91170b7 fix(trigger): surface webhook trigger quota exceeded as 429 response (#38656)
(cherry picked from commit 44fb074359)
2026-07-20 15:46:46 +08:00
林玮 (Jade Lin) 7be4546083 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>
(cherry picked from commit f93a5b95c6)
2026-07-20 15:46:46 +08:00
非法操作 2766439c7c fix: can't create chat app when get tenant default model schema error (#39266)
(cherry picked from commit 30df1433cc)
2026-07-20 15:12:54 +08:00
林玮 (Jade Lin) 3cc26b84bb refactor(api): pass tenant_id explicitly to workflow repositories (#39042)
(cherry picked from commit 678ce2ab37)
2026-07-20 14:42:56 +08:00
Bond ZhuandBond Zhu b3dd8ffcfc chore(api): add breadcrumb logs around vector backend factory resolution
Some production chat requests intermittently park forever between the
vector_db whitelist query and TidbOnQdrantVectorFactory.init_vector's
first log line: the message row stays status=normal with an empty answer,
no error is recorded, and no timeout ever fires. Live captures place the
freeze inside Vector._init_vector's factory resolution, but the API runs
gevent workers and a parked greenlet is invisible to py-spy (default and
--native modes both show only idle hubs), so the exact parking statement
cannot be captured from outside the process.

INFO-level breadcrumbs added:
- vector_factory._init_vector: before/after get_vector_factory, with
  tenant_id + dataset_id for correlating a hung request with its logs
- vector_backend_registry: cache-miss marker, and before/after ep.load()
  with the entry point value (module about to be imported)

Decoding the log tail of the next hung request:
- "resolving..." absent            -> parked before resolution
- "resolving..." only, no miss log -> parked on the cache-hit path or
                                      inside logging itself (handler lock)
- cache miss, no "loading..."      -> parked in entry_points() scan
- "loading...", no "... loaded"    -> parked inside ep.load() import
                                      (import lock / gevent interaction)
- "resolved...", no init_vector:   -> parked at the factory call boundary

All lines fire at most once per Vector construction; the slow-path lines
fire only on per-process first resolution, so added volume matches the
existing init_vector breadcrumb. Diagnostic instrumentation; no logic
changes, cheap to revert.
2026-07-20 14:21:09 +08:00
yyhand非法操作 c949d1da27 ci: route default Depot jobs to 4-core runners (#39277)
(cherry picked from commit 58d1dd8873)
2026-07-20 14:17:51 +08:00
非法操作 1328644dfc fix: change email not work (#39182)
(cherry picked from commit aff8a49bbf)
2026-07-20 14:17:51 +08:00
非法操作 2467df55c7 fix: can't debug model plugins (#38500)
(cherry picked from commit 6edce14e88)
2026-07-20 14:17:51 +08:00
非法操作 fed5a9a627 fix(web): redirect imported apps with creator permissions (#38460)
(cherry picked from commit 1247fa28f1)
2026-07-20 14:17:51 +08:00
林玮 (Jade Lin) 0a10d02a9f refactor(api): pass session explicitly through recommended app service chain 2026-07-15 15:57:49 +08:00
Xiyuan Chenand林玮 (Jade Lin) 5c0a0759f0 fix: preserve ResponseStreamFilter state across workflow pause/resume (#38540)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
(cherry picked from commit d72ee32ba1)
2026-07-15 14:47:22 +08:00
林玮 (Jade Lin) 3443054d96 fix(api): harden trial apps access for recommended apps (#38970)
(cherry picked from commit edca982b1e)
2026-07-15 14:36:35 +08:00
wangxiaoleiandXiyuan Chen be737a3504 feat: when knowledge permission is all team need update rbac setttings (#38911)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
(cherry picked from commit c278148c8f)
2026-07-14 03:52:24 -07:00
wangxiaoleiandXiyuan Chen e0b910d43d feat: create app init rbac scope to all and add workspace members. (#38636)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
(cherry picked from commit 7e05f28a46)
2026-07-14 03:52:24 -07:00
Coding On Star 44048320bd fix(auth): preserve redirect URL across auth flows (#38905)
Co-authored-by: CodingOnStar <hanxujiang@dify.com>
(cherry picked from commit 2c4d8ef098)
2026-07-14 18:23:54 +08:00
Coding On Star 3dc85bd398 fix(auth): prevent open redirects in post-login flows (#38864)
Co-authored-by: CodingOnStar <hanxujiang@dify.com>
(cherry picked from commit c68e5e5ed3)
2026-07-14 18:23:54 +08:00
SnakinyaandCoding On Star 56d87d6318 fix: SSRF bypass via raw httpx.get in API tool schema fetch (#37746)
Co-authored-by: 归青 <guiqing.lzw@antgroup.com>
Co-authored-by: QuantumGhost <obelisk.reg+git@gmail.com>
(cherry picked from commit ae0d6ee214)
2026-07-14 18:23:54 +08:00
林玮 (Jade Lin)andCoding On Star 73db4d3932 feat(oauth): preserve redirect_url through OAuth state for post-login… (#38900)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
(cherry picked from commit 82ff93cbdd)
2026-07-14 18:23:54 +08:00
林玮 (Jade Lin)andCoding On Star 13c6556e6f fix(api): require trial app registration for reads (#38858)
(cherry picked from commit 8e0d10bf5e)
2026-07-14 18:23:54 +08:00
Xin ZhangandXiyuan Chen 7570427f51 feat(enterprise): enforce licensed seats cap at account creation (#38883)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
(cherry picked from commit 28a036066a)
2026-07-14 00:22:05 -07:00
Xiyuan Chen 48717e5981 chore(deps): bump python-engineio, python-socketio, soupsieve for CVE fixes (#38844)
(cherry picked from commit c21a248df6)
2026-07-13 01:04:45 -07:00
Xiyuan Chen 5805b94e23 fix(api): wire dedicated timeout into inner RBAC requests (#38150)
(cherry picked from commit 0d5f43520c)
2026-07-12 23:39:30 -07:00
Jashwanth Reddy GummulaandXiyuan Chen c64c772b5b fix: resolve 36288 mypy errors (#37850)
Co-authored-by: WH-2099 <wh2099@pm.me>
(cherry picked from commit dd0c4a2296)
2026-07-12 21:53:31 -07:00
FFXNandXiyuan Chen b4537455cf fix: enhance SQL query safety and add metadata key validation (#38307)
(cherry picked from commit 88864ae086)
2026-07-12 21:53:31 -07:00
FFXNandXiyuan Chen 5fe010b653 fix: sql injection (#38295)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com>
(cherry picked from commit d9884efaee)
2026-07-12 21:53:31 -07:00
蒜香飞行矮堇瓜andXiyuan Chen 66cd69b07e fix(api): avoid infinite loop in _delete_records when batch deletion fails (#38118)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
(cherry picked from commit 089c3f4af0)
2026-07-12 21:53:31 -07:00
linhongkuanandXiyuan Chen 09eca00370 fix(api): isolate side-effect database writes (#37895)
Co-authored-by: FFXN <31929997+FFXN@users.noreply.github.com>
(cherry picked from commit d8f3be4bcd)
2026-07-12 21:53:31 -07:00
Xiyuan Chen 12cf69e5c8 fix(api): register rbac-migrate-dataset-permissions CLI command (#38204)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
(cherry picked from commit 5ce13d1773)
2026-07-12 21:53:31 -07:00
wangxiaoleiandXiyuan Chen 4332023806 feat: support dataset permission migrate to rbac (#38166)
(cherry picked from commit 528bf95d1b)
2026-07-12 21:53:31 -07:00
Wu TianweiandXiyuan Chen dd9f1e041e fix: update documentation links in permission set and role modals (#38188)
(cherry picked from commit 46163c68bd)
2026-07-12 21:53:31 -07:00
Wu TianweiandXiyuan Chen 2271257c2f chore(i18n): update role management permission keys for multiple languages (#38186)
(cherry picked from commit 54a8ff2c1d)
2026-07-12 21:53:31 -07:00
wangxiaoleiandXiyuan Chen 8adbda988b perf: make command rbac-migrate-member-roles use less mem and make it… (#38151)
(cherry picked from commit 49a92f096f)
2026-07-12 21:53:31 -07:00
-LAN-andXiyuan Chen 876ac24496 fix(api): require edit access for trace config changes (#37973)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
(cherry picked from commit 0035d90e36)
2026-07-12 21:53:31 -07:00
wangxiaolei 8b0892fd79 fix: fix auth prefix duplicate (#38616) 2026-07-10 10:51:05 +08:00
Coding On StarandCodingOnStar 50d2ad8479 fix: chunk workflow failure tracking data (#38598)
(cherry picked from commit 2f99652203)

Co-authored-by: CodingOnStar <hanxujiang@dify.com>
2026-07-09 17:11:46 +08:00
Crazywoolaandyyh 1ef8d3a2a9 fix: display errors for oauth page (#38546)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
(cherry picked from commit 5ffd4345dc)
2026-07-08 16:31:31 +08:00
yyh b8adb0cac2 refactor(web): move app context layout styles to shell (#38511)
(cherry picked from commit 2c6ec1a761)
2026-07-08 16:31:31 +08:00
yyh df88c45e3c chore(dify-ui): update theme tokens (#38189)
(cherry picked from commit f8afa49ab9)
2026-07-08 16:31:31 +08:00
autofix-ci[bot]andYansong Zhang d3df1b2847 [autofix.ci] apply automated fixes 2026-07-08 09:55:47 +08:00
Yansong Zhang 709aac9337 test: cover credit pool quota paths 2026-07-08 09:55:47 +08:00
Yansong Zhang 865affdd9b test: update quota billing service mocks 2026-07-08 09:55:47 +08:00
Yansong Zhang 71840c3f71 fix(api): route quota calls to new billing 2026-07-08 09:55:47 +08:00
Yansong Zhang dc09aea7b3 feat(api): use billing quota for credit pool 2026-07-08 09:55:47 +08:00
非法操作 99507a5ae5 fix: editor should not manage member (#38503)
(cherry picked from commit faaa4708a6)
2026-07-07 16:09:47 +08:00
55ae0aadf8 fix(api): resolve plugin external user ids in backwards invocation (#38098)
Co-authored-by: Harsh Kashyap <Harsh23Kashyap@users.noreply.github.com>
Co-authored-by: Harsh Kashyap <harshkashyap@Harshs-MacBook-Pro.local>
2026-07-06 17:12:39 +08:00
非法操作 6c4ba71a23 chore: use zstd for plugin model provider cache (#38382)
(cherry picked from commit 91cc50371b)
2026-07-03 17:19:43 +08:00
非法操作 e5f53053e3 fix: incorrect backgroud when dark mode and no plugins (#38378)
Co-authored-by: yyh <yuanyouhuilyz@gmail.com>
(cherry picked from commit 81cc43b753)
2026-07-03 17:19:43 +08:00
非法操作 3e392527e9 chore: compress large plugin_model obj in redis (#38374)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
(cherry picked from commit 059a1fe090)
2026-07-03 17:19:43 +08:00
WH-2099and非法操作 36d3a9c7fb fix(api): keep provider refresh waiters single-flight (#38226)
(cherry picked from commit 7fb797c121)
2026-07-03 17:19:43 +08:00
Jingyiand非法操作 914598d7e3 fix: prevent app card meta overflow (#38349)
(cherry picked from commit d120995efc)
2026-07-03 10:55:44 +08:00
yyhand非法操作 602c24cf8c fix(web): align main nav app item states (#38326)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: copilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com>
Co-authored-by: Jingyi-Dify <jingyi.qi@dify.ai>
(cherry picked from commit 57f687fb45)
2026-07-03 10:55:44 +08:00
非法操作 e0b5ecc50f fix: Notion sync empty state width in knowledge creation (#38321)
(cherry picked from commit 3d10f6c510)
2026-07-03 10:55:44 +08:00
Manan Bansaland非法操作 a83c7e89e0 fix(api): tolerate legacy service_api end-user type on load (#38271)
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
(cherry picked from commit 41fb484491)
2026-07-03 10:55:44 +08:00
yyhand非法操作 7a4252b3de fix(web): fill dataset create layout height (#38308)
(cherry picked from commit 6ec56e7bd1)
2026-07-02 16:25:42 +08:00
hjlarry 5c3dd25b32 fix: repair main nav hotfix backport
(cherry picked from commit 0bd359bb8f)
2026-07-02 15:25:06 +08:00
wangxiaoleiand非法操作 30c8e9b08d fix: Working outside of application context. (#38300)
(cherry picked from commit 826b259d9f)
2026-07-02 15:07:12 +08:00
Joeland非法操作 3e3a16feaf fix: web detail adjustment before release (#38296)
Co-authored-by: Jingyi-Dify <jingyi.qi@dify.ai>
(cherry picked from commit 0bd359bb8f)
2026-07-02 15:07:12 +08:00
hjlarry 0e1fad88c0 fix: avoid profile query in public web app chat
(cherry picked from commit 52c106b532)
2026-07-01 14:05:02 +08:00
yyhand非法操作 f442559543 fix: prevent exiting toasts from blocking page clicks (#38063)
(cherry picked from commit 484633d261)
2026-07-01 14:05:02 +08:00
Bond Zhuand非法操作 2a63f26796 fix(workflow): guard on_tool_execution stdout traces behind DEBUG (#38200)
(cherry picked from commit 03d59dba47)
2026-07-01 14:05:02 +08:00
非法操作 7e325b2ff2 fix: debug plugin permission setting not work (#38197)
(cherry picked from commit 4303103304)
2026-07-01 14:05:02 +08:00
Jingyiand非法操作 26487d19b4 fix: handle integration marketplace install callbacks (#38236)
Co-authored-by: hjlarry <hjlarry@163.com>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
(cherry picked from commit f816ae2e95)
2026-07-01 14:05:02 +08:00
非法操作 71a936e89a fix: editor can view the logs (#38165)
(cherry picked from commit fa1ac75922)
2026-07-01 14:05:02 +08:00
非法操作 f2c19b456b fix: editor should not query billing subscriptions (#38157)
(cherry picked from commit 34f62e7df6)
2026-07-01 14:05:02 +08:00
Manan Bansaland非法操作 2d3da2a49f fix(web): keep HITL input-field save button visible in the edit dialog (#38007)
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
(cherry picked from commit 267b34caaf)
2026-07-01 14:05:02 +08:00
非法操作 11b393db9f fix: agent app log detail modal not display well (#38014)
(cherry picked from commit 23917c7b3e)
2026-07-01 14:05:02 +08:00
非法操作 9a01b16d6c fix: plugin installation task popover layout when some failed too long (#38000)
(cherry picked from commit e22fd9efd6)
2026-07-01 14:05:02 +08:00
409 changed files with 12932 additions and 3902 deletions
+3 -3
View File
@@ -16,7 +16,7 @@ concurrency:
jobs:
api-unit:
name: API Unit Tests
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
env:
COVERAGE_FILE: coverage-unit
defaults:
@@ -75,7 +75,7 @@ jobs:
api-integration:
name: API Integration Tests
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
env:
COVERAGE_FILE: coverage-integration
STORAGE_TYPE: opendal
@@ -129,7 +129,7 @@ jobs:
api-coverage:
name: API Coverage
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
needs:
- api-unit
- api-integration
+1 -1
View File
@@ -171,7 +171,7 @@ jobs:
create-manifest:
needs: build
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
if: github.repository == 'langgenius/dify'
strategy:
matrix:
+2 -2
View File
@@ -23,7 +23,7 @@ concurrency:
jobs:
validate:
name: validate manifest + resolve target Dify release
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
if: github.repository == 'langgenius/dify'
permissions:
contents: read
@@ -87,7 +87,7 @@ jobs:
release:
name: build + attach standalone binaries (all targets)
needs: validate
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
permissions:
contents: write
defaults:
+2 -2
View File
@@ -9,7 +9,7 @@ concurrency:
jobs:
db-migration-test-postgres:
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Checkout code
@@ -59,7 +59,7 @@ jobs:
run: uv run --directory api flask upgrade-db
db-migration-test-mysql:
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Checkout code
+1 -1
View File
@@ -13,7 +13,7 @@ on:
jobs:
deploy:
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
if: |
github.event.workflow_run.conclusion == 'success' &&
github.event.workflow_run.head_branch == 'deploy/agent'
+1 -1
View File
@@ -10,7 +10,7 @@ on:
jobs:
deploy:
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
if: |
github.event.workflow_run.conclusion == 'success' &&
github.event.workflow_run.head_branch == 'deploy/dev'
+1 -1
View File
@@ -13,7 +13,7 @@ on:
jobs:
deploy:
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
if: |
github.event.workflow_run.conclusion == 'success' &&
github.event.workflow_run.head_branch == 'deploy/enterprise'
+1 -1
View File
@@ -13,7 +13,7 @@ on:
jobs:
deploy:
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
if: |
github.event.workflow_run.conclusion == 'success' &&
github.event.workflow_run.head_branch == 'deploy/saas'
+1 -1
View File
@@ -22,7 +22,7 @@ concurrency:
jobs:
check-cherry-pick-provenance:
name: Require cherry-pick provenance
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- uses: actions/checkout@df4cb1c069e1874edd31b4311f1884172cec0e10 # v6.0.3
with:
+1 -1
View File
@@ -7,7 +7,7 @@ jobs:
permissions:
contents: read
pull-requests: write
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- uses: actions/labeler@f27b608878404679385c85cfa523b85ccb86e213 # v6.1.0
with:
+14 -14
View File
@@ -23,7 +23,7 @@ concurrency:
jobs:
pre_job:
name: Skip Duplicate Checks
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
outputs:
should_skip: ${{ steps.skip_check.outputs.should_skip || 'false' }}
steps:
@@ -39,7 +39,7 @@ jobs:
name: Check Changed Files
needs: pre_job
if: needs.pre_job.outputs.should_skip != 'true'
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
outputs:
api-changed: ${{ steps.changes.outputs.api }}
cli-changed: ${{ steps.changes.outputs.cli }}
@@ -152,7 +152,7 @@ jobs:
- pre_job
- check-changes
if: needs.pre_job.outputs.should_skip != 'true' && needs.check-changes.outputs.api-changed != 'true'
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Report skipped API tests
run: echo "No API-related changes detected; skipping API tests."
@@ -165,7 +165,7 @@ jobs:
- check-changes
- api-tests-run
- api-tests-skip
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Finalize API Tests status
env:
@@ -212,7 +212,7 @@ jobs:
- pre_job
- check-changes
if: needs.pre_job.outputs.should_skip != 'true' && needs.check-changes.outputs.cli-changed != 'true'
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Report skipped CLI tests
run: echo "No CLI-related changes detected; skipping CLI tests."
@@ -225,7 +225,7 @@ jobs:
- check-changes
- cli-tests-run
- cli-tests-skip
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Finalize CLI Tests status
env:
@@ -272,7 +272,7 @@ jobs:
- pre_job
- check-changes
if: needs.pre_job.outputs.should_skip != 'true' && needs.check-changes.outputs.web-changed != 'true'
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Report skipped web tests
run: echo "No web-related changes detected; skipping web tests."
@@ -285,7 +285,7 @@ jobs:
- check-changes
- web-tests-run
- web-tests-skip
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Finalize Web Tests status
env:
@@ -331,7 +331,7 @@ jobs:
- pre_job
- check-changes
if: needs.pre_job.outputs.should_skip != 'true' && needs.check-changes.outputs.e2e-changed != 'true'
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Report skipped web full-stack e2e
run: echo "No E2E-related changes detected; skipping web full-stack E2E."
@@ -344,7 +344,7 @@ jobs:
- check-changes
- web-e2e-run
- web-e2e-skip
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Finalize Web Full-Stack E2E status
env:
@@ -396,7 +396,7 @@ jobs:
- pre_job
- check-changes
if: needs.pre_job.outputs.should_skip != 'true' && needs.check-changes.outputs.vdb-changed != 'true'
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Report skipped VDB tests
run: echo "No VDB-related changes detected; skipping VDB tests."
@@ -409,7 +409,7 @@ jobs:
- check-changes
- vdb-tests-run
- vdb-tests-skip
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Finalize VDB Tests status
env:
@@ -455,7 +455,7 @@ jobs:
- pre_job
- check-changes
if: needs.pre_job.outputs.should_skip != 'true' && needs.check-changes.outputs.migration-changed != 'true'
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Report skipped DB migration tests
run: echo "No migration-related changes detected; skipping DB migration tests."
@@ -468,7 +468,7 @@ jobs:
- check-changes
- db-migration-test-run
- db-migration-test-skip
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Finalize DB Migration Test status
env:
+1 -1
View File
@@ -12,7 +12,7 @@ permissions: {}
jobs:
comment:
name: Comment PR with pyrefly diff
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
permissions:
actions: read
contents: read
+1 -1
View File
@@ -10,7 +10,7 @@ permissions:
jobs:
pyrefly-diff:
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
permissions:
contents: read
issues: write
@@ -12,7 +12,7 @@ permissions: {}
jobs:
comment:
name: Comment PR with type coverage
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
permissions:
actions: read
contents: read
+1 -1
View File
@@ -10,7 +10,7 @@ permissions:
jobs:
pyrefly-type-coverage:
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
permissions:
contents: read
issues: write
+1 -1
View File
@@ -16,7 +16,7 @@ jobs:
name: Validate PR title
permissions:
pull-requests: read
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Complete merge group check
if: github.event_name == 'merge_group'
+1 -1
View File
@@ -12,7 +12,7 @@ on:
jobs:
stale:
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
permissions:
issues: write
pull-requests: write
+3 -3
View File
@@ -15,7 +15,7 @@ permissions:
jobs:
python-style:
name: Python Style
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Checkout code
@@ -61,7 +61,7 @@ jobs:
web-style:
name: Web Style
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
defaults:
run:
working-directory: ./web
@@ -175,7 +175,7 @@ jobs:
superlinter:
name: SuperLinter
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
steps:
- name: Checkout code
+1 -1
View File
@@ -17,7 +17,7 @@ concurrency:
jobs:
build:
name: unit test for Node.js SDK
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
defaults:
run:
+1 -1
View File
@@ -35,7 +35,7 @@ concurrency:
jobs:
translate:
if: github.repository == 'langgenius/dify'
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
timeout-minutes: 120
steps:
+1 -1
View File
@@ -16,7 +16,7 @@ concurrency:
jobs:
trigger:
if: github.repository == 'langgenius/dify'
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
timeout-minutes: 5
steps:
+1 -1
View File
@@ -16,7 +16,7 @@ jobs:
test:
name: Full VDB Tests
if: github.repository == 'langgenius/dify'
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
strategy:
matrix:
python-version:
+1 -1
View File
@@ -13,7 +13,7 @@ concurrency:
jobs:
test:
name: VDB Smoke Tests
runs-on: depot-ubuntu-24.04
runs-on: depot-ubuntu-24.04-4
strategy:
matrix:
python-version:
+20 -2
View File
@@ -2,14 +2,14 @@ import logging
import time
import socketio
from flask import request
from flask import g, request
from opentelemetry.trace import get_current_span
from opentelemetry.trace.span import INVALID_SPAN_ID, INVALID_TRACE_ID
from configs import dify_config
from contexts.wrapper import RecyclableContextVar
from controllers.console.error import UnauthorizedAndForceLogout
from core.logging.context import init_request_context
from core.logging.context import clear_request_context, init_request_context
from dify_app import DifyApp
from extensions.ext_socketio import sio
from services.enterprise.enterprise_service import EnterpriseService
@@ -60,6 +60,7 @@ def create_flask_app_with_configs() -> DifyApp:
def before_request():
# Initialize logging context for this request
init_request_context()
g._logging_context_cleanup_scheduled = False
RecyclableContextVar.increment_thread_recycles()
# Enterprise license validation for API endpoints (both console and webapp)
@@ -117,9 +118,26 @@ def create_flask_app_with_configs() -> DifyApp:
logger.warning("Failed to add trace headers to response", exc_info=True)
return response
@dify_app.after_request
def schedule_logging_context_cleanup(response):
"""Keep logging context through streaming, then clear it when WSGI closes the response."""
if response.direct_passthrough:
return response
g._logging_context_cleanup_scheduled = True
response.call_on_close(clear_request_context)
return response
@dify_app.teardown_request
def clear_unscheduled_logging_context(_error: BaseException | None) -> None:
"""Clear when no response-close callback can own cleanup."""
if not g.get("_logging_context_cleanup_scheduled", False):
clear_request_context()
# Capture the decorator return values so static checkers do not treat the hooks as unused.
_ = before_request
_ = add_trace_headers
_ = schedule_logging_context_cleanup
_ = clear_unscheduled_logging_context
return dify_app
+2 -1
View File
@@ -22,7 +22,7 @@ from .plugin import (
setup_system_trigger_oauth_client,
transform_datasource_credentials,
)
from .rbac import migrate_member_roles_to_rbac
from .rbac import migrate_dataset_permissions_to_rbac, migrate_member_roles_to_rbac
from .retention import (
archive_workflow_runs,
archive_workflow_runs_plan,
@@ -76,6 +76,7 @@ __all__ = [
"legacy_model_types",
"migrate_annotation_vector_database",
"migrate_data_for_plugin",
"migrate_dataset_permissions_to_rbac",
"migrate_knowledge_vector_database",
"migrate_member_roles_to_rbac",
"migrate_oss",
+2
View File
@@ -7,6 +7,7 @@ from typing import cast
import click
from commands.rbac import migrate_dataset_permissions_to_rbac
from extensions.ext_database import db
from graphon.model_runtime.entities.model_entities import ModelType
from services.legacy_model_type_migration import (
@@ -177,3 +178,4 @@ def legacy_model_types(
data_migrate.add_command(legacy_model_types)
data_migrate.add_command(migrate_dataset_permissions_to_rbac)
+438 -66
View File
@@ -1,11 +1,55 @@
from __future__ import annotations
import json
from collections.abc import Iterator
from concurrent.futures import ThreadPoolExecutor, as_completed
import click
from sqlalchemy import select
from configs import dify_config
from core.db.session_factory import session_factory
from models import TenantAccountJoin, TenantAccountRole
from services.enterprise.rbac_service import ListOption, RBACService
from core.rbac import RBACResourceWhitelistScope
from models import Dataset, DatasetPermission, DatasetPermissionEnum, TenantAccountJoin, TenantAccountRole
from services.enterprise.rbac_service import ListOption, RBACService, ReplaceMemberBindings, ReplaceUserAccessPolicies
_RBAC_DEFAULT_ACCESS_POLICY_ID = "default"
_LEGACY_ROLE_TO_BUILTIN_TAG = {
TenantAccountRole.OWNER.value: "owner",
TenantAccountRole.ADMIN.value: "admin",
TenantAccountRole.EDITOR.value: "editor",
TenantAccountRole.NORMAL.value: "normal",
TenantAccountRole.DATASET_OPERATOR.value: "dataset_operator",
}
def _resolve_builtin_role_ids(tenant_id: str, operator_account_id: str) -> dict[str, str]:
"""Resolve every legacy workspace role to the current tenant's builtin RBAC role id.
The migration replays the old `TenantAccountJoin.role` values onto the
RBAC member-role binding API. Builtin RBAC roles are tenant-scoped and
identified by runtime ids, so the command must look them up per tenant.
"""
roles = RBACService.Roles.list(
tenant_id=tenant_id,
account_id=operator_account_id,
options=ListOption(page_number=1, results_per_page=100),
).data
role_id_by_tag = {
role.role_tag: role.id
for role in roles
if role.is_builtin and role.category == "global_system_default" and role.role_tag
}
resolved: dict[str, str] = {}
for legacy_role, expected_builtin_tag in _LEGACY_ROLE_TO_BUILTIN_TAG.items():
role_id = role_id_by_tag.get(expected_builtin_tag)
if expected_builtin_tag == "dataset_operator" and not dify_config.DATASET_OPERATOR_ENABLED:
continue
if not role_id:
raise ValueError(f"Builtin RBAC role not found for tenant={tenant_id}, legacy_role={legacy_role}")
resolved[legacy_role] = role_id
return resolved
def _resolve_builtin_role_id(tenant_id: str, operator_account_id: str, legacy_role: str) -> str:
@@ -15,26 +59,86 @@ def _resolve_builtin_role_id(tenant_id: str, operator_account_id: str, legacy_ro
RBAC member-role binding API. Builtin RBAC roles are tenant-scoped and
identified by runtime ids, so the command must look them up per tenant.
"""
expected_builtin_tag = {
TenantAccountRole.OWNER.value: "owner",
TenantAccountRole.ADMIN.value: "admin",
TenantAccountRole.EDITOR.value: "editor",
TenantAccountRole.NORMAL.value: "normal",
TenantAccountRole.DATASET_OPERATOR.value: "dataset_operator",
}.get(legacy_role)
if not expected_builtin_tag:
if legacy_role not in _LEGACY_ROLE_TO_BUILTIN_TAG:
raise ValueError(f"Unsupported legacy workspace role: {legacy_role}")
roles = RBACService.Roles.list(
return _resolve_builtin_role_ids(tenant_id, operator_account_id)[legacy_role]
def _iter_tenant_member_batches(
tenant_id: str | None,
*,
db_batch_size: int,
api_batch_size: int,
) -> Iterator[tuple[str, str, list[tuple[str, str]]]]:
"""Yield legacy member roles in tenant-scoped API-sized batches.
Rows are projected to primitive values and streamed from the database, so
the command never materializes every TenantAccountJoin ORM object. The
iterator only keeps one tenant's API-sized batches in memory while it
finds that tenant's owner account.
"""
with session_factory.create_session() as session:
stmt = (
select(TenantAccountJoin.tenant_id, TenantAccountJoin.account_id, TenantAccountJoin.role)
.order_by(TenantAccountJoin.tenant_id.asc(), TenantAccountJoin.id.asc())
.execution_options(yield_per=db_batch_size)
)
if tenant_id:
stmt = stmt.where(TenantAccountJoin.tenant_id == tenant_id)
current_tenant_id: str | None = None
owner_account_id: str | None = None
batches: list[list[tuple[str, str]]] = []
batch: list[tuple[str, str]] = []
def flush_current_tenant() -> Iterator[tuple[str, str, list[tuple[str, str]]]]:
if current_tenant_id is None:
return
if batch:
batches.append(batch.copy())
if not owner_account_id:
raise ValueError(f"Workspace owner not found for tenant={current_tenant_id}")
for item in batches:
yield current_tenant_id, owner_account_id, item
for row in session.execute(stmt):
workspace_id = str(row.tenant_id)
if current_tenant_id is not None and workspace_id != current_tenant_id:
yield from flush_current_tenant()
owner_account_id = None
batches = []
batch = []
current_tenant_id = workspace_id
account_id = str(row.account_id)
role = str(row.role)
if role == TenantAccountRole.OWNER.value:
owner_account_id = account_id
batch.append((account_id, role))
if len(batch) >= api_batch_size:
batches.append(batch)
batch = []
yield from flush_current_tenant()
def _member_already_has_role(current_roles_by_account_id: dict[str, set[str]], account_id: str, role_id: str) -> bool:
return current_roles_by_account_id.get(account_id) == {role_id}
def _replace_member_role(
tenant_id: str,
operator_account_id: str,
member_account_id: str,
role_id: str,
) -> str:
RBACService.MemberRoles.replace(
tenant_id=tenant_id,
account_id=operator_account_id,
options=ListOption(page_number=1, results_per_page=100),
).data
for role in roles:
if role.is_builtin and role.category == "global_system_default" and role.role_tag == expected_builtin_tag:
return str(role.id)
raise ValueError(f"Builtin RBAC role not found for tenant={tenant_id}, legacy_role={legacy_role}")
member_account_id=member_account_id,
role_ids=[role_id],
)
return member_account_id
@click.command(
@@ -42,7 +146,16 @@ def _resolve_builtin_role_id(tenant_id: str, operator_account_id: str, legacy_ro
)
@click.option("--tenant-id", help="Only migrate a single workspace.")
@click.option("--dry-run", is_flag=True, default=False, help="Preview the migration without writing RBAC bindings.")
def migrate_member_roles_to_rbac(tenant_id: str | None, dry_run: bool) -> None:
@click.option("--db-batch-size", default=5000, show_default=True, help="Rows fetched per database batch.")
@click.option("--api-batch-size", default=200, show_default=True, help="Members checked per RBAC batch_get call.")
@click.option("--workers", default=1, show_default=True, help="Concurrent member role replace calls per tenant batch.")
def migrate_member_roles_to_rbac(
tenant_id: str | None,
dry_run: bool,
db_batch_size: int,
api_batch_size: int,
workers: int,
) -> None:
"""Backfill RBAC member-role bindings from legacy `TenantAccountJoin.role` data.
This is an offline migration command for workspaces that already have
@@ -50,63 +163,322 @@ def migrate_member_roles_to_rbac(tenant_id: str | None, dry_run: bool) -> None:
member-role binding store.
"""
click.echo(click.style("Starting RBAC member-role migration.", fg="green"))
if workers < 1:
raise click.BadParameter("workers must be >= 1", param_hint="--workers")
with session_factory.create_session() as session:
stmt = select(TenantAccountJoin).order_by(TenantAccountJoin.tenant_id.asc(), TenantAccountJoin.id.asc())
if tenant_id:
stmt = stmt.where(TenantAccountJoin.tenant_id == tenant_id)
tenant_count = 0
scanned_count = 0
skipped_count = 0
migrated_count = 0
current_tenant_id: str | None = None
role_ids_by_legacy_role: dict[str, str] = {}
joins = list(session.scalars(stmt).all())
for workspace_id, owner_account_id, batch in _iter_tenant_member_batches(
tenant_id,
db_batch_size=db_batch_size,
api_batch_size=api_batch_size,
):
scanned_count += len(batch)
if workspace_id != current_tenant_id:
tenant_count += 1
current_tenant_id = workspace_id
role_ids_by_legacy_role = _resolve_builtin_role_ids(workspace_id, owner_account_id)
click.echo(f"tenant={workspace_id}")
if not joins:
current_roles_by_account_id: dict[str, set[str]] = {}
if not dry_run:
current_roles = RBACService.MemberRoles.batch_get(
tenant_id=workspace_id,
account_id=owner_account_id,
member_account_ids=[account_id for account_id, _ in batch],
)
current_roles_by_account_id = {item.account_id: {role.id for role in item.roles} for item in current_roles}
replace_jobs: list[tuple[str, str]] = []
for member_account_id, legacy_role in batch:
resolved_role_id = role_ids_by_legacy_role.get(legacy_role)
if not resolved_role_id:
raise ValueError(f"Unsupported legacy workspace role: {legacy_role}")
if dry_run:
click.echo(
f"tenant={workspace_id} member={member_account_id} "
f"legacy_role={legacy_role} -> rbac_role_id={resolved_role_id}"
)
continue
if _member_already_has_role(current_roles_by_account_id, member_account_id, resolved_role_id):
skipped_count += 1
continue
replace_jobs.append((member_account_id, resolved_role_id))
if replace_jobs:
if workers == 1:
for member_account_id, resolved_role_id in replace_jobs:
_replace_member_role(workspace_id, owner_account_id, member_account_id, resolved_role_id)
migrated_count += 1
else:
with ThreadPoolExecutor(max_workers=workers) as executor:
futures = [
executor.submit(
_replace_member_role,
workspace_id,
owner_account_id,
member_account_id,
resolved_role_id,
)
for member_account_id, resolved_role_id in replace_jobs
]
for future in as_completed(futures):
future.result()
migrated_count += 1
if scanned_count % 10000 == 0:
click.echo(
f"progress scanned={scanned_count} migrated={migrated_count} skipped={skipped_count}",
err=True,
)
if scanned_count == 0:
click.echo(click.style("No workspace members found for migration.", fg="yellow"))
return
owner_account_by_tenant: dict[str, str] = {}
resolved_role_ids: dict[tuple[str, str], str] = {}
migrated_count = 0
for join in joins:
workspace_id = str(join.tenant_id)
member_account_id = str(join.account_id)
legacy_role = str(join.role)
if workspace_id not in owner_account_by_tenant:
owner_join = next(
(
item
for item in joins
if str(item.tenant_id) == workspace_id and str(item.role) == TenantAccountRole.OWNER.value
),
None,
)
if not owner_join:
raise ValueError(f"Workspace owner not found for tenant={workspace_id}")
owner_account_by_tenant[workspace_id] = str(owner_join.account_id)
operator_account_id = owner_account_by_tenant[workspace_id]
cache_key = (workspace_id, legacy_role)
if cache_key not in resolved_role_ids:
resolved_role_ids[cache_key] = _resolve_builtin_role_id(workspace_id, operator_account_id, legacy_role)
resolved_role_id = resolved_role_ids[cache_key]
if dry_run:
click.echo(
f"tenant={workspace_id} member={member_account_id} "
f"legacy_role={legacy_role} -> rbac_role_id={resolved_role_id}"
click.style(
f"Dry run completed. Scanned {scanned_count} members across {tenant_count} tenants. "
"No RBAC bindings were written.",
fg="yellow",
)
)
else:
click.echo(
click.style(
f"RBAC member-role migration completed. Scanned {scanned_count} members across {tenant_count} tenants, "
f"migrated {migrated_count}, skipped {skipped_count} already up-to-date.",
fg="green",
)
)
if dry_run:
continue
RBACService.MemberRoles.replace(
tenant_id=workspace_id,
account_id=operator_account_id,
member_account_id=member_account_id,
role_ids=[resolved_role_id],
)
migrated_count += 1
def _dataset_permission_enum(permission: DatasetPermissionEnum | str | None) -> DatasetPermissionEnum:
if permission is None:
return DatasetPermissionEnum.ONLY_ME
try:
return DatasetPermissionEnum(permission)
except ValueError as exc:
raise ValueError(f"Unsupported legacy dataset permission: {permission}") from exc
def _rbac_dataset_scope_for_legacy_permission(permission: DatasetPermissionEnum) -> RBACResourceWhitelistScope:
if permission is DatasetPermissionEnum.ALL_TEAM:
return RBACResourceWhitelistScope.ALL
if permission in {DatasetPermissionEnum.ONLY_ME, DatasetPermissionEnum.PARTIAL_TEAM}:
return RBACResourceWhitelistScope.SPECIFIC
raise ValueError(f"Unsupported legacy dataset permission: {permission}")
def _emit_dataset_permission_migration_event(payload: dict[str, object]) -> None:
click.echo(json.dumps(payload, sort_keys=True))
@click.command(
"rbac-migrate-dataset-permissions",
help=(
"Migrate legacy dataset permission scopes and partial members into RBAC dataset access bindings. "
"Side effect: replacing each dataset whitelist clears existing per-user policy bindings; "
"the command then recreates legacy partial-member default bindings."
),
)
@click.option("--tenant-id", help="Only migrate datasets in a single workspace.")
@click.option("--dataset-id", help="Only migrate a single dataset.")
@click.option("--batch-size", default=500, show_default=True, type=click.IntRange(min=1))
@click.option(
"--dry-run/--apply",
default=True,
show_default=True,
help="Preview the migration without writing RBAC bindings. Use --apply to write changes.",
)
def migrate_dataset_permissions_to_rbac(
tenant_id: str | None,
dataset_id: str | None,
batch_size: int,
dry_run: bool,
) -> None:
"""Backfill RBAC dataset access config from legacy `Dataset.permission`.
Legacy mapping:
- all_team_members -> RBAC dataset whitelist scope "all"
- partial_members -> RBAC dataset whitelist scope "specific" plus each partial member gets the
virtual default policy
- only_me -> RBAC dataset whitelist scope "specific" with no member policy bindings
The command replaces each dataset's RBAC whitelist scope first. RBAC clears
existing per-user policy bindings during that replace, then this command
recreates the legacy partial-member default bindings. Re-running it is
therefore idempotent for a dataset's current legacy configuration.
"""
click.echo(click.style("Starting RBAC dataset permission migration.", fg="green"))
scanned_count = 0
scope_migrated_count = 0
user_policy_migrated_count = 0
partial_dataset_count = 0
last_dataset_id: str | None = None
while True:
with session_factory.create_session() as session:
stmt = (
select(Dataset.id, Dataset.tenant_id, Dataset.permission, Dataset.created_by)
.order_by(Dataset.id.asc())
.limit(batch_size)
)
if tenant_id:
stmt = stmt.where(Dataset.tenant_id == tenant_id)
if dataset_id:
stmt = stmt.where(Dataset.id == dataset_id)
if last_dataset_id:
stmt = stmt.where(Dataset.id > last_dataset_id)
dataset_rows = list(session.execute(stmt).all())
if not dataset_rows:
break
dataset_ids = [str(row.id) for row in dataset_rows]
partial_members_by_dataset_id: dict[str, list[str]] = {item: [] for item in dataset_ids}
permission_rows = session.execute(
select(DatasetPermission.dataset_id, DatasetPermission.account_id).where(
DatasetPermission.dataset_id.in_(dataset_ids)
)
).all()
for row in permission_rows:
partial_members_by_dataset_id[str(row.dataset_id)].append(str(row.account_id))
for dataset in dataset_rows:
workspace_id = str(dataset.tenant_id)
current_dataset_id = str(dataset.id)
operator_account_id = str(dataset.created_by)
permission_value = _dataset_permission_enum(dataset.permission)
scope = _rbac_dataset_scope_for_legacy_permission(permission_value)
partial_member_ids = sorted(set(partial_members_by_dataset_id[current_dataset_id]))
should_bind_partial_members = permission_value is DatasetPermissionEnum.PARTIAL_TEAM
click.echo(
f"tenant={workspace_id} dataset={current_dataset_id} "
f"operator={operator_account_id} "
f"legacy_permission={permission_value} -> rbac_scope={scope} "
f"partial_members={len(partial_member_ids) if should_bind_partial_members else 0}"
)
scanned_count += 1
replace_whitelist_payload = ReplaceMemberBindings(scope=scope)
if dry_run:
_emit_dataset_permission_migration_event(
{
"event": "dataset_permission_migration_proposed_change",
"action": "replace_whitelist",
"dry_run": True,
"tenant_id": workspace_id,
"dataset_id": current_dataset_id,
"operator_account_id": operator_account_id,
"before": {
"legacy_dataset_permission": permission_value.value,
"legacy_partial_member_ids": partial_member_ids if should_bind_partial_members else [],
},
"after": {
"rbac_whitelist_scope": scope.value,
},
"call": {
"method": "RBACService.DatasetAccess.replace_whitelist",
"kwargs": {
"tenant_id": workspace_id,
"account_id": operator_account_id,
"dataset_id": current_dataset_id,
"payload": replace_whitelist_payload.model_dump(mode="json"),
},
},
}
)
if not dry_run:
RBACService.DatasetAccess.replace_whitelist(
tenant_id=workspace_id,
account_id=operator_account_id,
dataset_id=current_dataset_id,
payload=replace_whitelist_payload,
)
scope_migrated_count += 1
if should_bind_partial_members:
partial_dataset_count += 1
for member_account_id in partial_member_ids:
replace_user_access_policies_payload = ReplaceUserAccessPolicies(
access_policy_ids=[_RBAC_DEFAULT_ACCESS_POLICY_ID],
)
if dry_run:
_emit_dataset_permission_migration_event(
{
"event": "dataset_permission_migration_proposed_change",
"action": "replace_user_access_policies",
"dry_run": True,
"tenant_id": workspace_id,
"dataset_id": current_dataset_id,
"operator_account_id": operator_account_id,
"target_account_id": member_account_id,
"before": {
"legacy_dataset_permission": permission_value.value,
"legacy_partial_member_id": member_account_id,
},
"after": {
"rbac_user_access_policy_ids": [_RBAC_DEFAULT_ACCESS_POLICY_ID],
},
"call": {
"method": "RBACService.DatasetAccess.replace_user_access_policies",
"kwargs": {
"tenant_id": workspace_id,
"account_id": operator_account_id,
"dataset_id": current_dataset_id,
"target_account_id": member_account_id,
"payload": replace_user_access_policies_payload.model_dump(
mode="json", exclude_unset=True
),
},
},
}
)
continue
RBACService.DatasetAccess.replace_user_access_policies(
tenant_id=workspace_id,
account_id=operator_account_id,
dataset_id=current_dataset_id,
target_account_id=member_account_id,
payload=replace_user_access_policies_payload,
)
user_policy_migrated_count += 1
last_dataset_id = dataset_ids[-1]
if dataset_id:
break
if scanned_count == 0:
click.echo(click.style("No datasets found for migration.", fg="yellow"))
return
if dry_run:
click.echo(click.style("Dry run completed. No RBAC bindings were written.", fg="yellow"))
click.echo(
click.style(
f"Dry run completed. Scanned {scanned_count} datasets; "
f"{partial_dataset_count} partial-member datasets would be migrated.",
fg="yellow",
)
)
else:
click.echo(click.style(f"RBAC member-role migration completed. Migrated {migrated_count} members.", fg="green"))
click.echo(
click.style(
"RBAC dataset permission migration completed. "
f"Scanned {scanned_count} datasets, migrated {scope_migrated_count} scopes, "
f"wrote {user_policy_migrated_count} user default-policy bindings.",
fg="green",
)
)
+6
View File
@@ -34,6 +34,12 @@ class EnterpriseFeatureConfig(BaseSettings):
default=False,
)
ENTERPRISE_RBAC_REQUEST_TIMEOUT: int = Field(
ge=1,
description="Maximum timeout in seconds for inner RBAC requests.",
default=30,
)
class EnterpriseTelemetryConfig(BaseSettings):
"""
+10
View File
@@ -43,6 +43,7 @@ from controllers.console.wraps import (
from core.ops.ops_trace_manager import OpsTraceManager
from core.rag.entities import PreProcessingRule, Rule, Segmentation
from core.rag.retrieval.retrieval_methods import RetrievalMethod
from core.rbac import RBACResourceWhitelistScope
from core.trigger.constants import TRIGGER_NODE_TYPES
from extensions.ext_database import db
from fields.base import ResponseModel
@@ -69,6 +70,7 @@ from services.entities.knowledge_entities.knowledge_entities import (
WeightVectorSetting,
)
from services.feature_service import FeatureService
from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task
ALLOW_CREATE_APP_MODES = ["chat", "agent-chat", "advanced-chat", "workflow", "completion"]
@@ -666,6 +668,14 @@ class AppListApi(Resource):
app_service = AppService()
app = app_service.create_app(current_tenant_id, params, current_user)
if dify_config.RBAC_ENABLED:
enterprise_rbac_service.RBACService.AppAccess.replace_whitelist(
tenant_id=str(current_tenant_id),
account_id=current_user.id,
app_id=str(app.id),
payload=enterprise_rbac_service.ReplaceMemberBindings(scope=RBACResourceWhitelistScope.ALL),
)
initialize_created_app_rbac_access_task.delay(current_tenant_id, current_user.id, app_id=app.id)
permission_keys_map = enterprise_rbac_service.RBACService.AppPermissions.batch_get(
str(current_tenant_id),
current_user.id,
+10
View File
@@ -13,6 +13,7 @@ from controllers.console.wraps import (
RBACPermission,
RBACResourceScope,
account_initialization_required,
edit_permission_required,
rbac_permission_required,
setup_required,
)
@@ -95,9 +96,12 @@ class TraceAppConfigApi(Resource):
console_ns.models[TraceAppConfigResponse.__name__],
)
@console_ns.response(400, "Invalid request parameters or configuration already exists")
@console_ns.response(403, "Insufficient permissions")
@setup_required
@login_required
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TRACING_CONFIG)
@get_app_model
def post(self, app_model: App):
"""Create a new trace app configuration"""
@@ -125,9 +129,12 @@ class TraceAppConfigApi(Resource):
console_ns.models[TraceAppConfigResponse.__name__],
)
@console_ns.response(400, "Invalid request parameters or configuration not found")
@console_ns.response(403, "Insufficient permissions")
@setup_required
@login_required
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TRACING_CONFIG)
@get_app_model
def patch(self, app_model: App):
"""Update an existing trace app configuration"""
@@ -149,9 +156,12 @@ class TraceAppConfigApi(Resource):
@console_ns.doc(params=query_params_from_model(TraceProviderQuery))
@console_ns.response(204, "Tracing configuration deleted successfully")
@console_ns.response(400, "Invalid request parameters or configuration not found")
@console_ns.response(403, "Insufficient permissions")
@setup_required
@login_required
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TRACING_CONFIG)
@get_app_model
def delete(self, app_model: App):
"""Delete an existing trace app configuration"""
+10 -2
View File
@@ -19,7 +19,8 @@ from controllers.console.app.error import AppNotFoundError
from core.db.session_factory import session_factory
from extensions.ext_database import db
from libs.login import current_account_with_tenant
from models import App, AppMode
from models import App, AppMode, TrialApp
from services.recommended_app_service import RecommendedAppService
def _load_app_model(session: Session, app_id: str) -> App | None:
@@ -41,7 +42,10 @@ def _load_app_model_from_scoped_session(app_id: str) -> App | None:
def _load_app_model_with_trial(app_id: str) -> App | None:
app_model = db.session.scalar(select(App).where(App.id == app_id, App.status == "normal").limit(1))
"""Load a normal app through its trial registration without applying current-tenant scope."""
app_model = db.session.scalar(
select(App).join(TrialApp, TrialApp.app_id == App.id).where(App.id == app_id, App.status == "normal").limit(1)
)
return app_model
@@ -195,6 +199,8 @@ def get_app_model_with_trial[**P, R](
*,
mode: AppMode | list[AppMode] | None = None,
) -> Callable[P, R] | Callable[[Callable[P, R]], Callable[P, R]]:
"""Inject an app registered for trial or available from the recommended catalog."""
def decorator(view_func: Callable[P, R]) -> Callable[P, R]:
@wraps(view_func)
def decorated_view(*args: P.args, **kwargs: P.kwargs) -> R:
@@ -207,6 +213,8 @@ def get_app_model_with_trial[**P, R](
del kwargs["app_id"]
app_model = _load_app_model_with_trial(app_id)
if app_model is None:
app_model = RecommendedAppService.get_app(app_id, session=db.session())
if not app_model:
raise AppNotFoundError()
@@ -23,9 +23,9 @@ from libs.password import valid_password
from models import Account
from services.account_service import AccountService
from services.billing_service import BillingService
from services.errors.account import AccountRegisterError
from services.errors.account import AccountRegisterError, SeatsLimitExceededError
from ..error import AccountInFreezeError, EmailSendIpLimitError
from ..error import AccountInFreezeError, EmailSendIpLimitError, SeatsLimitExceeded
from ..wraps import email_password_login_enabled, email_register_enabled, setup_required
@@ -208,5 +208,7 @@ class EmailRegisterResetApi(Resource):
timezone=timezone,
session=db.session,
)
except SeatsLimitExceededError:
raise SeatsLimitExceeded()
except AccountRegisterError:
raise AccountInFreezeError()
+4 -1
View File
@@ -25,6 +25,7 @@ from controllers.console.error import (
AccountNotFound,
EmailSendIpLimitError,
NotAllowedCreateWorkspace,
SeatsLimitExceeded,
WorkspacesLimitExceeded,
)
from controllers.console.wraps import (
@@ -51,7 +52,7 @@ from models.account import Account
from services.account_service import AccountService, InvitationDetailDict, RegisterService, TenantService
from services.billing_service import BillingService
from services.entities.auth_entities import LoginFailureReason, LoginPayloadBase
from services.errors.account import AccountRegisterError
from services.errors.account import AccountRegisterError, SeatsLimitExceededError
from services.errors.workspace import WorkSpaceNotAllowedCreateError, WorkspacesLimitExceededError
from services.feature_service import FeatureService
@@ -317,6 +318,8 @@ class EmailCodeLoginApi(Resource):
)
except WorkSpaceNotAllowedCreateError:
raise NotAllowedCreateWorkspace()
except SeatsLimitExceededError:
raise SeatsLimitExceeded()
except AccountRegisterError:
_log_console_login_failure(email=user_email, reason=LoginFailureReason.ACCOUNT_IN_FREEZE)
raise AccountInFreezeError()
+40 -4
View File
@@ -25,7 +25,7 @@ from libs.token import (
from models import Account, AccountStatus
from services.account_service import AccountService, RegisterService, TenantService
from services.billing_service import BillingService
from services.errors.account import AccountNotFoundError, AccountRegisterError
from services.errors.account import AccountNotFoundError, AccountRegisterError, SeatsLimitExceededError
from services.errors.workspace import WorkSpaceNotAllowedCreateError, WorkSpaceNotFoundError
from services.feature_service import FeatureService
@@ -38,6 +38,7 @@ class OAuthLoginQuery(BaseModel):
invite_token: str | None = Field(default=None, description="Optional invitation token")
timezone: str | None = Field(default=None, description="Preferred timezone")
language: str | None = Field(default=None, description="Preferred interface language")
redirect_url: str | None = Field(default=None, description="Relative page to resume after login")
class OAuthCallbackQuery(BaseModel):
@@ -87,6 +88,36 @@ def _validated_language(value: str | None) -> str | None:
return None
def _url_origin(url: str) -> tuple[str, str, int] | None:
parsed_url = urllib.parse.urlsplit(url)
if parsed_url.scheme not in {"http", "https"} or parsed_url.hostname is None:
return None
try:
port = parsed_url.port
except ValueError:
return None
if port is None:
port = 443 if parsed_url.scheme == "https" else 80
return parsed_url.scheme, parsed_url.hostname, port
def _get_redirect_target(redirect_url: str | None) -> str:
if not redirect_url:
return dify_config.CONSOLE_WEB_URL
parsed_url = urllib.parse.urlsplit(redirect_url)
normalized_path = redirect_url.lstrip().replace("\\", "/")
if not parsed_url.scheme and not parsed_url.netloc and not normalized_path.startswith("//"):
return redirect_url
redirect_origin = _url_origin(redirect_url)
if redirect_origin is not None and redirect_origin == _url_origin(dify_config.CONSOLE_WEB_URL):
return redirect_url
return dify_config.CONSOLE_WEB_URL
def _preferred_interface_language(language: str | None = None) -> str:
if language:
return language
@@ -109,6 +140,7 @@ class OAuthLogin(Resource):
invite_token = request.args.get("invite_token") or None
timezone = _validated_timezone(request.args.get("timezone") or None)
language = _validated_language(request.args.get("language") or None)
redirect_url = request.args.get("redirect_url") or None
OAUTH_PROVIDERS = get_oauth_providers()
with current_app.app_context():
oauth_provider = OAUTH_PROVIDERS.get(provider)
@@ -119,6 +151,7 @@ class OAuthLogin(Resource):
invite_token=invite_token,
timezone=timezone,
language=language,
redirect_url=redirect_url,
)
return redirect(auth_url)
@@ -144,6 +177,7 @@ class OAuthCallback(Resource):
invite_token = oauth_state.get("invite_token")
timezone = _validated_timezone(oauth_state.get("timezone"))
language = _validated_language(oauth_state.get("language"))
redirect_url = oauth_state.get("redirect_url")
if not code:
return {"error": "Authorization code is required"}, 400
@@ -182,6 +216,8 @@ class OAuthCallback(Resource):
f"{dify_config.CONSOLE_WEB_URL}/signin"
"?message=Workspace not found, please contact system admin to invite you to join in a workspace."
)
except SeatsLimitExceededError:
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message=Licensed seats limit exceeded.")
except AccountRegisterError as e:
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message={e.description}")
@@ -210,9 +246,9 @@ class OAuthCallback(Resource):
ip_address=extract_remote_ip(request),
)
base_url = dify_config.CONSOLE_WEB_URL
query_char = "&" if "?" in base_url else "?"
target_url = f"{base_url}{query_char}oauth_new_user={str(oauth_new_user).lower()}"
target_url = _get_redirect_target(redirect_url)
query_char = "&" if "?" in target_url else "?"
target_url = f"{target_url}{query_char}oauth_new_user={str(oauth_new_user).lower()}"
response = redirect(target_url)
set_access_token_to_cookie(request, response, token_pair.access_token)
@@ -50,6 +50,8 @@ from models.provider_ids import ModelProviderID
from services.api_token_service import ApiTokenCache
from services.dataset_service import DatasetPermissionService, DatasetService, DocumentService
from services.enterprise import rbac_service as enterprise_rbac_service
from services.enterprise.rbac_service import RBACResourceWhitelistScope, ReplaceMemberBindings
from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task
register_response_schema_models(console_ns, ApiBaseUrlResponse, SimpleResultResponse, UsageCheckResponse)
@@ -565,6 +567,16 @@ class DatasetListApi(Resource):
except services.errors.dataset.DatasetNameDuplicateError:
raise DatasetNameDuplicateError()
if dify_config.RBAC_ENABLED:
if permission == DatasetPermissionEnum.ALL_TEAM:
enterprise_rbac_service.RBACService.DatasetAccess.replace_whitelist(
current_tenant_id,
current_user.id,
dataset.id,
ReplaceMemberBindings(scope=RBACResourceWhitelistScope.ALL),
)
initialize_created_app_rbac_access_task.delay(current_tenant_id, current_user.id, dataset_id=dataset.id)
permission_keys_map = enterprise_rbac_service.RBACService.DatasetPermissions.batch_get(
str(current_tenant_id),
current_user.id,
+6
View File
@@ -58,6 +58,12 @@ class WorkspacesLimitExceeded(BaseHTTPException):
code = 400
class SeatsLimitExceeded(BaseHTTPException):
error_code = "limit_exceeded"
description = "Unable to create account because the licensed seats limit was exceeded"
code = 400
class AccountBannedError(BaseHTTPException):
error_code = "account_banned"
description = "Account is banned."
@@ -106,7 +106,7 @@ class RecommendedAppListApi(Resource):
language_prefix = _resolve_language(args.language, current_user)
return RecommendedAppListResponse.model_validate(
RecommendedAppService.get_recommended_apps_and_categories(db.session, language_prefix),
RecommendedAppService.get_recommended_apps_and_categories(language_prefix, session=db.session()),
from_attributes=True,
).model_dump(mode="json")
+48 -3
View File
@@ -44,7 +44,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
@@ -64,6 +66,7 @@ from fields.app_fields import (
tag_fields,
)
from fields.dataset_fields import dataset_fields
from fields.file_fields import FileResponse, FileWithSignedUrl
from fields.member_fields import simple_account_fields
from fields.message_fields import SuggestedQuestionsResponse
from fields.workflow_fields import (
@@ -75,8 +78,8 @@ from fields.workflow_fields import (
from graphon.graph_engine.manager import GraphEngineManager
from graphon.model_runtime.errors.invoke import InvokeError
from libs import helper
from libs.helper import uuid_value
from models import Account
from libs.helper import dump_response, uuid_value
from models import Account, App
from models.account import TenantStatus
from models.model import AppMode, Site
from models.workflow import Workflow
@@ -199,6 +202,36 @@ register_response_schema_models(
)
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__])
@@ -605,6 +638,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",
+56 -33
View File
@@ -28,7 +28,7 @@ from controllers.console.wraps import (
from extensions.ext_database import db
from fields.file_fields import FileResponse, UploadConfig
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
@@ -38,6 +38,59 @@ register_response_schema_models(console_ns, AllowedExtensionsResponse, TextConte
PREVIEW_WORDS_LIMIT = 3000
FILE_UPLOAD_PARAMS = {
"file": {
"description": "File to upload",
"in": "formData",
"type": "file",
"required": True,
},
"source": {
"description": "Optional upload source",
"in": "formData",
"type": "string",
"enum": ["datasets"],
"required": False,
},
}
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):
@@ -64,41 +117,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.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)
response = FileResponse.model_validate(upload_file, from_attributes=True)
return response.model_dump(mode="json"), 201
+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,
)
+9 -1
View File
@@ -36,7 +36,7 @@ from libs.login import current_account_with_tenant, login_required
from models.account import Account, TenantAccountJoin, TenantAccountRole
from services.account_service import AccountService, RegisterService, TenantService
from services.enterprise import rbac_service as enterprise_rbac_service
from services.errors.account import AccountAlreadyInTenantError
from services.errors.account import AccountAlreadyInTenantError, SeatsLimitExceededError
from services.feature_service import FeatureService
@@ -291,6 +291,14 @@ class MemberInviteEmailApi(Resource):
"message": "Account already in workspace.",
}
)
except SeatsLimitExceededError:
invitation_results.append(
{
"status": "failed",
"email": invitee_email,
"message": "Licensed seats limit exceeded.",
}
)
except Exception as e:
invitation_results.append({"status": "failed", "email": invitee_email, "message": str(e)})
+13 -17
View File
@@ -1,6 +1,5 @@
from __future__ import annotations
from enum import StrEnum
from typing import Any
from flask import request
@@ -14,9 +13,11 @@ from controllers.common.schema import register_response_schema_models
from controllers.console import console_ns
from controllers.console.wraps import RBACPermission, RBACResourceScope, rbac_permission_required
from core.db.session_factory import session_factory
from core.rbac import RBACResourceWhitelistScope
from libs.login import current_account_with_tenant, login_required
from models import Account
from services.enterprise import rbac_service as svc
from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task
class _RBACRoleList(svc.Paginated[svc.RBACRole]):
@@ -511,14 +512,8 @@ class RBACAccessPolicyBindingUnlockApi(Resource):
# ---------------------------------------------------------------------------
class _AccessScope(StrEnum):
ALL = "all"
SPECIFIC = "specific"
ONLY_ME = "only_me"
class _ResourceAccessScopeRequest(BaseModel):
scope: _AccessScope
scope: RBACResourceWhitelistScope
class _ReplaceBindingsRequest(BaseModel):
@@ -584,14 +579,15 @@ class RBACAppWhitelistApi(Resource):
def put(self, app_id):
tenant_id, account_id = _current_ids()
request = _payload(_ResourceAccessScopeRequest)
return _dump(
svc.RBACService.AppAccess.replace_whitelist(
tenant_id,
account_id,
str(app_id),
svc.ReplaceMemberBindings(scope=request.scope.value),
)
result = svc.RBACService.AppAccess.replace_whitelist(
tenant_id,
account_id,
str(app_id),
svc.ReplaceMemberBindings(scope=request.scope.value),
)
if dify_config.RBAC_ENABLED and request.scope is RBACResourceWhitelistScope.ALL:
initialize_created_app_rbac_access_task.delay(tenant_id, account_id, str(app_id))
return _dump(result)
@console_ns.route("/workspaces/current/rbac/apps/<uuid:app_id>/user-access-policies")
@@ -616,8 +612,8 @@ class RBACAppUserAccessPolicyAssignmentApi(Resource):
svc.RBACService.AppAccess.replace_user_access_policies(
tenant_id,
account_id,
str(app_id),
str(target_account_id),
app_id,
target_account_id,
payload,
)
)
+3
View File
@@ -48,6 +48,7 @@ from services.errors.account import (
MemberNotInTenantError,
NoPermissionError,
RoleAlreadyAssignedError,
SeatsLimitExceededError,
)
from services.feature_service import FeatureService
@@ -190,6 +191,8 @@ class WorkspaceMembersApi(Resource):
raise BadRequest(str(exc))
except NoPermissionError as exc:
raise BadRequest(str(exc))
except SeatsLimitExceededError:
raise BadRequest("licensed seats limit exceeded")
except AccountRegisterError as exc:
raise BadRequest(str(exc))
@@ -43,6 +43,7 @@ from models.dataset import DatasetPermissionEnum
from models.enums import TagType
from models.provider_ids import ModelProviderID
from services.dataset_service import DatasetPermissionService, DatasetService, DocumentService
from services.enterprise.rbac_service import RBACResourceWhitelistScope, RBACService, ReplaceMemberBindings
from services.entities.knowledge_entities.knowledge_entities import (
ExternalRetrievalModel,
KnowledgeProvider,
@@ -58,6 +59,7 @@ from services.tag_service import (
from services.tag_service import (
UpdateTagPayload as UpdateTagServicePayload,
)
from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task
register_enum_models(service_api_ns, DatasetPermissionEnum)
@@ -523,6 +525,15 @@ 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)
return _dump_service_dataset_detail(dataset), 200
+3 -2
View File
@@ -12,6 +12,7 @@ from flask_restx import Resource
from flask_restx.utils import merge
from pydantic import BaseModel
from sqlalchemy import select
from sqlalchemy.orm import sessionmaker
from werkzeug.exceptions import Forbidden, NotFound, Unauthorized
from configs import dify_config
@@ -269,8 +270,8 @@ def cloud_edition_billing_rate_limit_check[**P, R](
subscription_plan=knowledge_rate_limit.subscription_plan,
operation="knowledge",
)
db.session.add(rate_limit_log)
db.session.commit()
with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session:
session.add(rate_limit_log)
raise Forbidden(
"Sorry, you have reached the knowledge base request rate limit of your subscription."
)
+25 -2
View File
@@ -7,10 +7,26 @@ from werkzeug.exceptions import NotFound, RequestEntityTooLarge
from controllers.trigger import bp
from core.trigger.debug.event_bus import TriggerDebugEventBus
from core.trigger.debug.events import WebhookDebugEvent, build_webhook_pool_key
from enums.quota_type import QuotaType
from services.errors.app import QuotaExceededError
from services.trigger.webhook_service import RawWebhookDataDict, WebhookService
logger = logging.getLogger(__name__)
_QUOTA_EXCEEDED_MESSAGES = {
QuotaType.TRIGGER: "Trigger event quota exceeded. Please upgrade your plan.",
QuotaType.WORKFLOW: "Workflow execution quota exceeded. Please upgrade your plan.",
}
_DEFAULT_QUOTA_EXCEEDED_MESSAGE = "Quota exceeded. Please upgrade your plan."
def _get_quota_exceeded_message(feature: str) -> str:
try:
quota_type = QuotaType(feature)
except ValueError:
return _DEFAULT_QUOTA_EXCEEDED_MESSAGE
return _QUOTA_EXCEEDED_MESSAGES.get(quota_type, _DEFAULT_QUOTA_EXCEEDED_MESSAGE)
def _prepare_webhook_execution(webhook_id: str, is_debug: bool = False):
"""Fetch trigger context, extract request data, and validate payload using unified processing.
@@ -60,8 +76,15 @@ def handle_webhook(webhook_id: str):
response_data, status_code = WebhookService.generate_webhook_response(node_config)
return jsonify(response_data), status_code
except ValueError as e:
raise NotFound(str(e))
except QuotaExceededError as error:
return jsonify(
{
"error": "Too Many Requests",
"message": _get_quota_exceeded_message(error.feature),
}
), 429
except ValueError as error:
raise NotFound(str(error))
except RequestEntityTooLarge:
raise
except Exception as e:
+7 -1
View File
@@ -29,6 +29,7 @@ from controllers.console.wraps import (
)
from controllers.web import web_ns
from controllers.web.wraps import decode_jwt_token
from core.logging.context import set_identity_context
from libs.helper import EmailStr, extract_remote_ip
from libs.passport import PassportService
from libs.password import valid_password
@@ -160,7 +161,12 @@ class LoginStatusApi(Resource):
user_logged_in = False
try:
_ = decode_jwt_token(app_code=app_code, user_id=user_id)
_, end_user = decode_jwt_token(app_code=app_code, user_id=user_id)
set_identity_context(
tenant_id=end_user.tenant_id,
user_id=end_user.id,
user_type=end_user.type or "end_user",
)
app_logged_in = True
except Exception:
app_logged_in = False
+6
View File
@@ -11,6 +11,7 @@ from werkzeug.exceptions import BadRequest, NotFound, Unauthorized
from constants import HEADER_NAME_APP_CODE
from controllers.web.error import WebAppAuthAccessDeniedError, WebAppAuthRequiredError
from core.logging.context import set_identity_context
from extensions.ext_database import db
from libs.passport import PassportService
from libs.token import extract_webapp_passport
@@ -28,6 +29,11 @@ def validate_jwt_token[**P, R](
@wraps(view)
def decorated(*args: P.args, **kwargs: P.kwargs) -> R:
app_model, end_user = decode_jwt_token()
set_identity_context(
tenant_id=end_user.tenant_id,
user_id=end_user.id,
user_type=end_user.type or "end_user",
)
return view(app_model, end_user, *args, **kwargs)
return decorated
@@ -47,6 +47,7 @@ from core.repositories import DifyCoreRepositoryFactory
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
from extensions.ext_database import db
from factories import file_factory
from graphon.filters import ResponseStreamFilter
from graphon.graph_engine.layers import GraphEngineLayer
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
from graphon.runtime import GraphRuntimeState
@@ -231,6 +232,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
workflow_triggered_from = WorkflowRunTriggeredFrom.APP_RUN
workflow_execution_repository = DifyCoreRepositoryFactory.create_workflow_execution_repository(
session_factory=session_factory,
tenant_id=app_model.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=workflow_triggered_from,
@@ -238,6 +240,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
# Create workflow node execution repository
workflow_node_execution_repository = DifyCoreRepositoryFactory.create_workflow_node_execution_repository(
session_factory=session_factory,
tenant_id=app_model.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN,
@@ -268,6 +271,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
graph_runtime_state: GraphRuntimeState,
pause_state_config: PauseStateLayerConfig | None = None,
response_stream_filter: ResponseStreamFilter | None = None,
):
"""
Resume a paused advanced chat execution.
@@ -297,6 +301,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
stream=application_generate_entity.stream,
pause_state_config=pause_state_config,
graph_runtime_state=graph_runtime_state,
response_stream_filter=response_stream_filter,
)
def single_iteration_generate(
@@ -356,6 +361,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
# Create workflow execution(aka workflow run) repository
workflow_execution_repository = DifyCoreRepositoryFactory.create_workflow_execution_repository(
session_factory=session_factory,
tenant_id=app_model.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowRunTriggeredFrom.DEBUGGING,
@@ -363,6 +369,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
# Create workflow node execution repository
workflow_node_execution_repository = DifyCoreRepositoryFactory.create_workflow_node_execution_repository(
session_factory=session_factory,
tenant_id=app_model.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP,
@@ -443,6 +450,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
# Create workflow execution(aka workflow run) repository
workflow_execution_repository = DifyCoreRepositoryFactory.create_workflow_execution_repository(
session_factory=session_factory,
tenant_id=app_model.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowRunTriggeredFrom.DEBUGGING,
@@ -450,6 +458,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
# Create workflow node execution repository
workflow_node_execution_repository = DifyCoreRepositoryFactory.create_workflow_node_execution_repository(
session_factory=session_factory,
tenant_id=app_model.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP,
@@ -491,6 +500,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
pause_state_config: PauseStateLayerConfig | None = None,
graph_runtime_state: GraphRuntimeState | None = None,
graph_engine_layers: Sequence[GraphEngineLayer] = (),
response_stream_filter: ResponseStreamFilter | None = None,
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
"""
Generate App response.
@@ -538,12 +548,14 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
)
graph_layers: list[GraphEngineLayer] = list(graph_engine_layers)
resolved_response_stream_filter = response_stream_filter or ResponseStreamFilter()
if pause_state_config is not None:
graph_layers.append(
PauseStatePersistenceLayer(
session_factory=pause_state_config.session_factory,
generate_entity=application_generate_entity,
state_owner_user_id=pause_state_config.state_owner_user_id,
response_stream_filter=resolved_response_stream_filter,
)
)
@@ -564,6 +576,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
"workflow_node_execution_repository": workflow_node_execution_repository,
"graph_engine_layers": tuple(graph_layers),
"graph_runtime_state": graph_runtime_state,
"response_stream_filter": resolved_response_stream_filter,
},
)
@@ -585,7 +598,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)
@@ -603,6 +620,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
graph_engine_layers: Sequence[GraphEngineLayer] = (),
graph_runtime_state: GraphRuntimeState | None = None,
response_stream_filter: ResponseStreamFilter | None = None,
):
"""
Generate worker in a new thread.
@@ -662,6 +680,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
workflow_node_execution_repository=workflow_node_execution_repository,
graph_engine_layers=graph_engine_layers,
graph_runtime_state=graph_runtime_state,
response_stream_filter=response_stream_filter,
)
try:
@@ -39,6 +39,7 @@ from extensions.ext_database import db
from extensions.ext_redis import redis_client
from extensions.otel import WorkflowAppRunnerHandler, trace_span
from graphon.enums import WorkflowType
from graphon.filters import ResponseStreamFilter
from graphon.graph_engine.command_channels import RedisChannel
from graphon.graph_engine.layers import GraphEngineLayer
from graphon.runtime import GraphRuntimeState, VariablePool
@@ -73,6 +74,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
graph_engine_layers: Sequence[GraphEngineLayer] = (),
graph_runtime_state: GraphRuntimeState | None = None,
response_stream_filter: ResponseStreamFilter | None = None,
):
super().__init__(
queue_manager=queue_manager,
@@ -90,6 +92,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
self._workflow_execution_repository = workflow_execution_repository
self._workflow_node_execution_repository = workflow_node_execution_repository
self._resume_graph_runtime_state = graph_runtime_state
self._response_stream_filter = response_stream_filter
@trace_span(WorkflowAppRunnerHandler)
def run(self):
@@ -227,6 +230,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
variable_pool=variable_pool,
graph_runtime_state=graph_runtime_state,
command_channel=command_channel,
response_stream_filter=self._response_stream_filter,
)
self._queue_manager.graph_runtime_state = graph_runtime_state
+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,
@@ -214,6 +214,7 @@ class PipelineGenerator(BaseAppGenerator):
session_factory = sessionmaker(bind=db.engine, expire_on_commit=False)
workflow_execution_repository = DifyCoreRepositoryFactory.create_workflow_execution_repository(
session_factory=session_factory,
tenant_id=pipeline.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=workflow_triggered_from,
@@ -221,6 +222,7 @@ class PipelineGenerator(BaseAppGenerator):
workflow_node_execution_repository = DifyCoreRepositoryFactory.create_workflow_node_execution_repository(
session_factory=session_factory,
tenant_id=pipeline.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.RAG_PIPELINE_RUN,
@@ -341,6 +343,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(
@@ -417,6 +420,7 @@ class PipelineGenerator(BaseAppGenerator):
workflow_execution_repository = DifyCoreRepositoryFactory.create_workflow_execution_repository(
session_factory=session_factory,
tenant_id=pipeline.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowRunTriggeredFrom.RAG_PIPELINE_DEBUGGING,
@@ -424,6 +428,7 @@ class PipelineGenerator(BaseAppGenerator):
workflow_node_execution_repository = DifyCoreRepositoryFactory.create_workflow_node_execution_repository(
session_factory=session_factory,
tenant_id=pipeline.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP,
@@ -513,6 +518,7 @@ class PipelineGenerator(BaseAppGenerator):
workflow_execution_repository = DifyCoreRepositoryFactory.create_workflow_execution_repository(
session_factory=session_factory,
tenant_id=pipeline.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowRunTriggeredFrom.RAG_PIPELINE_DEBUGGING,
@@ -520,6 +526,7 @@ class PipelineGenerator(BaseAppGenerator):
workflow_node_execution_repository = DifyCoreRepositoryFactory.create_workflow_node_execution_repository(
session_factory=session_factory,
tenant_id=pipeline.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP,
+20 -1
View File
@@ -42,6 +42,7 @@ from core.repositories import DifyCoreRepositoryFactory
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
from extensions.ext_database import db
from factories import file_factory
from graphon.filters import ResponseStreamFilter
from graphon.graph_engine.layers import GraphEngineLayer
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
from graphon.runtime import GraphRuntimeState
@@ -241,6 +242,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
workflow_triggered_from = WorkflowRunTriggeredFrom.APP_RUN
workflow_execution_repository = DifyCoreRepositoryFactory.create_workflow_execution_repository(
session_factory=session_factory,
tenant_id=app_model.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=workflow_triggered_from,
@@ -248,6 +250,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
# Create workflow node execution repository
workflow_node_execution_repository = DifyCoreRepositoryFactory.create_workflow_node_execution_repository(
session_factory=session_factory,
tenant_id=app_model.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN,
@@ -280,6 +283,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
graph_engine_layers: Sequence[GraphEngineLayer] = (),
pause_state_config: PauseStateLayerConfig | None = None,
variable_loader: VariableLoader = DUMMY_VARIABLE_LOADER,
response_stream_filter: ResponseStreamFilter | None = None,
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
"""
Resume a paused workflow execution using the persisted runtime state.
@@ -310,6 +314,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
graph_engine_layers=graph_engine_layers,
graph_runtime_state=graph_runtime_state,
pause_state_config=pause_state_config,
response_stream_filter=response_stream_filter,
)
def _generate(
@@ -328,6 +333,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
graph_engine_layers: Sequence[GraphEngineLayer] = (),
graph_runtime_state: GraphRuntimeState | None = None,
pause_state_config: PauseStateLayerConfig | None = None,
response_stream_filter: ResponseStreamFilter | None = None,
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
"""
Generate App response.
@@ -356,12 +362,14 @@ class WorkflowAppGenerator(BaseAppGenerator):
app_mode=app_model.mode,
)
resolved_response_stream_filter = response_stream_filter or ResponseStreamFilter()
if pause_state_config is not None:
graph_layers.append(
PauseStatePersistenceLayer(
session_factory=pause_state_config.session_factory,
generate_entity=application_generate_entity,
state_owner_user_id=pause_state_config.state_owner_user_id,
response_stream_filter=resolved_response_stream_filter,
)
)
@@ -384,12 +392,17 @@ class WorkflowAppGenerator(BaseAppGenerator):
"workflow_node_execution_repository": workflow_node_execution_repository,
"graph_engine_layers": tuple(graph_layers),
"graph_runtime_state": graph_runtime_state,
"response_stream_filter": resolved_response_stream_filter,
},
)
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(
@@ -459,6 +472,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
# Create workflow execution(aka workflow run) repository
workflow_execution_repository = DifyCoreRepositoryFactory.create_workflow_execution_repository(
session_factory=session_factory,
tenant_id=app_model.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowRunTriggeredFrom.DEBUGGING,
@@ -466,6 +480,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
# Create workflow node execution repository
workflow_node_execution_repository = DifyCoreRepositoryFactory.create_workflow_node_execution_repository(
session_factory=session_factory,
tenant_id=app_model.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP,
@@ -546,6 +561,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
# Create workflow execution(aka workflow run) repository
workflow_execution_repository = DifyCoreRepositoryFactory.create_workflow_execution_repository(
session_factory=session_factory,
tenant_id=app_model.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowRunTriggeredFrom.DEBUGGING,
@@ -553,6 +569,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
# Create workflow node execution repository
workflow_node_execution_repository = DifyCoreRepositoryFactory.create_workflow_node_execution_repository(
session_factory=session_factory,
tenant_id=app_model.tenant_id,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP,
@@ -590,6 +607,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
root_node_id: str | None = None,
graph_engine_layers: Sequence[GraphEngineLayer] = (),
graph_runtime_state: GraphRuntimeState | None = None,
response_stream_filter: ResponseStreamFilter | None = None,
) -> None:
"""
Generate worker in a new thread.
@@ -638,6 +656,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
root_node_id=root_node_id,
graph_engine_layers=graph_engine_layers,
graph_runtime_state=graph_runtime_state,
response_stream_filter=response_stream_filter,
)
try:
+4
View File
@@ -18,6 +18,7 @@ from core.workflow.workflow_entry import WorkflowEntry
from extensions.ext_redis import redis_client
from extensions.otel import WorkflowAppRunnerHandler, trace_span
from graphon.enums import WorkflowType
from graphon.filters import ResponseStreamFilter
from graphon.graph_engine.command_channels import RedisChannel
from graphon.graph_engine.layers import GraphEngineLayer
from graphon.runtime import GraphRuntimeState, VariablePool
@@ -46,6 +47,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
graph_engine_layers: Sequence[GraphEngineLayer] = (),
graph_runtime_state: GraphRuntimeState | None = None,
response_stream_filter: ResponseStreamFilter | None = None,
):
super().__init__(
queue_manager=queue_manager,
@@ -60,6 +62,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
self._workflow_execution_repository = workflow_execution_repository
self._workflow_node_execution_repository = workflow_node_execution_repository
self._resume_graph_runtime_state = graph_runtime_state
self._response_stream_filter = response_stream_filter
@trace_span(WorkflowAppRunnerHandler)
def run(self):
@@ -163,6 +166,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
variable_pool=variable_pool,
graph_runtime_state=graph_runtime_state,
command_channel=command_channel,
response_stream_filter=self._response_stream_filter,
)
persistence_layer = WorkflowPersistenceLayer(
@@ -7,6 +7,7 @@ from sqlalchemy.orm import Session, sessionmaker
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, WorkflowAppGenerateEntity
from core.workflow.system_variables import SystemVariableKey, get_system_text
from graphon.filters import ResponseStreamFilter
from graphon.graph_engine.layers import GraphEngineLayer
from graphon.graph_events import GraphEngineEvent, GraphRunPausedEvent
from models.model import AppMode
@@ -41,6 +42,10 @@ class WorkflowResumptionContext(BaseModel):
# Only workflow / chatflow could be paused.
generate_entity: _GenerateEntityUnion
serialized_graph_runtime_state: str
# Optional so that a workflow run paused before this field existed still
# loads: it just degrades to fresh-filter behavior on resume for that one
# stale run.
serialized_response_stream_filter_state: str | None = None
def dumps(self) -> str:
return self.model_dump_json()
@@ -52,6 +57,12 @@ class WorkflowResumptionContext(BaseModel):
def get_generate_entity(self) -> WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity:
return self.generate_entity.entity
def get_response_stream_filter(self) -> ResponseStreamFilter:
response_stream_filter = ResponseStreamFilter()
if self.serialized_response_stream_filter_state is not None:
response_stream_filter.loads(self.serialized_response_stream_filter_state)
return response_stream_filter
@dataclass(frozen=True)
class PauseStateLayerConfig:
@@ -67,11 +78,17 @@ class PauseStatePersistenceLayer(GraphEngineLayer):
session_factory: Engine | sessionmaker[Session],
generate_entity: WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity,
state_owner_user_id: str,
response_stream_filter: ResponseStreamFilter,
):
"""Create a PauseStatePersistenceLayer.
The `state_owner_user_id` is used when creating state file for pause.
It generally should id of the creator of workflow.
`response_stream_filter` must be the exact same instance that
`WorkflowEntry` is using to stream this run's events — this layer
dumps its state on pause, and a different instance would silently
persist the wrong (empty) filter state.
"""
if isinstance(session_factory, Engine):
session_factory = sessionmaker(session_factory)
@@ -79,6 +96,7 @@ class PauseStatePersistenceLayer(GraphEngineLayer):
self._session_maker = session_factory
self._state_owner_user_id = state_owner_user_id
self._generate_entity = generate_entity
self._response_stream_filter = response_stream_filter
def _get_repo(self) -> APIWorkflowRunRepository:
return DifyAPIRepositoryFactory.create_api_workflow_run_repository(self._session_maker)
@@ -119,6 +137,7 @@ class PauseStatePersistenceLayer(GraphEngineLayer):
state = WorkflowResumptionContext(
serialized_graph_runtime_state=self.graph_runtime_state.dumps(),
generate_entity=entity_wrapper,
serialized_response_stream_filter_state=self._response_stream_filter.dumps(),
)
workflow_run_id = get_system_text(
@@ -1,6 +1,7 @@
from collections.abc import Generator, Iterable, Mapping
from typing import Any
from configs import dify_config
from core.callback_handler.agent_tool_callback_handler import DifyAgentCallbackHandler, print_text
from core.ops.ops_trace_manager import TraceQueueManager
from core.tools.entities.tool_entities import ToolInvokeMessage
@@ -19,8 +20,9 @@ class DifyWorkflowCallbackHandler(DifyAgentCallbackHandler):
trace_manager: TraceQueueManager | None = None,
) -> Generator[ToolInvokeMessage, None, None]:
for tool_output in tool_outputs:
print_text("\n[on_tool_execution]\n", color=self.color)
print_text("Tool: " + tool_name + "\n", color=self.color)
print_text("Outputs: " + tool_output.model_dump_json()[:1000] + "\n", color=self.color)
print_text("\n")
if dify_config.DEBUG:
print_text("\n[on_tool_execution]\n", color=self.color)
print_text("Tool: " + tool_name + "\n", color=self.color)
print_text("Outputs: " + tool_output.model_dump_json()[:1000] + "\n", color=self.color)
print_text("\n")
yield tool_output
@@ -13,7 +13,7 @@ from core.helper.code_executor.jinja2.jinja2_transformer import Jinja2TemplateTr
from core.helper.code_executor.python3.python3_transformer import Python3TemplateTransformer
from core.helper.code_executor.template_transformer import TemplateTransformer
from core.helper.http_client_pooling import get_pooled_http_client
from graphon.nodes.code.entities import CodeLanguage
from graphon.nodes.code.entities import CodeLanguage as CodeLanguage # noqa: PLC0414
logger = logging.getLogger(__name__)
code_execution_endpoint_url = URL(str(dify_config.CODE_EXECUTION_ENDPOINT))
@@ -133,7 +133,9 @@ class CodeExecutor:
return response_code.data.stdout or ""
@classmethod
def execute_workflow_code_template(cls, language: CodeLanguage, code: str, inputs: Mapping[str, Any]):
def execute_workflow_code_template(
cls, language: CodeLanguage, code: str, inputs: Mapping[str, Any]
) -> dict[str, Any]:
"""
Execute code
:param language: code language
@@ -11,7 +11,7 @@ class Jinja2TemplateTransformer(TemplateTransformer):
@classmethod
@override
def transform_response(cls, response: str):
def transform_response(cls, response: str) -> dict[str, Any]:
"""
Transform response to dict
:param response: response
@@ -36,14 +36,14 @@ class TemplateTransformer(ABC):
return runner_script, preload_script
@classmethod
def extract_result_str_from_response(cls, response: str):
def extract_result_str_from_response(cls, response: str) -> str:
result = re.search(rf"{cls._result_tag}(.*){cls._result_tag}", response, re.DOTALL)
if not result:
raise ValueError(f"Failed to parse result: no result tag found in response. Response: {response[:200]}...")
return result.group(1)
@classmethod
def transform_response(cls, response: str) -> Mapping[str, Any]:
def transform_response(cls, response: str) -> dict[str, Any]:
"""
Transform response to dict
:param response: response
@@ -71,7 +71,7 @@ class TemplateTransformer(ABC):
return result
@classmethod
def _post_process_result(cls, result: dict[Any, Any]) -> dict[Any, Any]:
def _post_process_result(cls, result: dict[str, Any]) -> dict[str, Any]:
"""
Post-process the result to convert scientific notation strings back to numbers
"""
@@ -89,7 +89,7 @@ class TemplateTransformer(ABC):
return [convert_scientific_notation(v) for v in value]
return value
return convert_scientific_notation(result)
return {key: convert_scientific_notation(value) for key, value in result.items()}
@classmethod
@abstractmethod
+1 -1
View File
@@ -24,7 +24,7 @@ def upload_dsl(dsl_file_bytes: bytes, filename: str = "template.yaml") -> str:
response.raise_for_status()
data = response.json()
claim_code = data.get("data", {}).get("claim_code")
if not claim_code:
if not isinstance(claim_code, str) or not claim_code:
raise ValueError("Creators Platform did not return a valid claim_code")
return claim_code
+7 -4
View File
@@ -10,18 +10,21 @@ def is_credential_exists(credential_id: str, credential_type: "PluginCredentialT
"""
Check if the credential still exists in the database.
Uses the configured SQLAlchemy session factory instead of Flask-SQLAlchemy's
``db.engine`` because workflow graph node construction may run without an
active Flask application context.
:param credential_id: The credential ID to check
:param credential_type: The type of credential (MODEL or TOOL)
:return: True if credential exists, False otherwise
"""
from sqlalchemy import select
from sqlalchemy.orm import Session
from extensions.ext_database import db
from core.db import session_factory
from models.provider import ProviderCredential, ProviderModelCredential
from models.tools import BuiltinToolProvider
with Session(db.engine) as session:
with session_factory.create_session() as session:
if credential_type == PluginCredentialType.MODEL:
# Check both pre-defined and custom model credentials using a single UNION query
stmt = (
@@ -42,7 +45,7 @@ def is_credential_exists(credential_id: str, credential_type: "PluginCredentialT
def runtime_check_credential_policy_compliance(
credential_id: str, provider: str, credential_type: "PluginCredentialType", check_existence: bool = True
):
) -> None:
if dify_config.ENTERPRISE_DISABLE_RUNTIME_CREDENTIAL_CHECK:
return
check_credential_policy_compliance(
+4 -1
View File
@@ -1,4 +1,7 @@
def download_with_size_limit(url, max_download_size: int, **kwargs):
from typing import Any
def download_with_size_limit(url: str, max_download_size: int, **kwargs: Any) -> bytes:
from core.file import remote_fetcher
response = remote_fetcher.make_request("GET", url, follow_redirects=True, **kwargs)
+8 -6
View File
@@ -1,5 +1,7 @@
import base64
from Crypto.PublicKey import RSA
from libs import rsa
@@ -11,13 +13,13 @@ def obfuscated_token(token: str) -> str:
return token[:6] + "*" * 12 + token[-2:]
def full_mask_token(token_length=20):
def full_mask_token(token_length: int = 20) -> str:
return "*" * token_length
def encrypt_token(tenant_id: str, token: str):
from extensions.ext_database import db
def encrypt_token(tenant_id: str, token: str) -> str:
from models.account import Tenant
from models.engine import db
if not (tenant := db.session.get(Tenant, tenant_id)):
raise ValueError(f"Tenant with id {tenant_id} not found")
@@ -30,15 +32,15 @@ def decrypt_token(tenant_id: str, token: str) -> str:
return rsa.decrypt(base64.b64decode(token), tenant_id)
def batch_decrypt_token(tenant_id: str, tokens: list[str]):
def batch_decrypt_token(tenant_id: str, tokens: list[str]) -> list[str]:
rsa_key, cipher_rsa = rsa.get_decrypt_decoding(tenant_id)
return [rsa.decrypt_token_with_decoding(base64.b64decode(token), rsa_key, cipher_rsa) for token in tokens]
def get_decrypt_decoding(tenant_id: str):
def get_decrypt_decoding(tenant_id: str) -> tuple[RSA.RsaKey, object]:
return rsa.get_decrypt_decoding(tenant_id)
def decrypt_token_with_decoding(token: str, rsa_key, cipher_rsa):
def decrypt_token_with_decoding(token: str, rsa_key: RSA.RsaKey, cipher_rsa: object) -> str:
return rsa.decrypt_token_with_decoding(base64.b64decode(token), rsa_key, cipher_rsa)
+14 -4
View File
@@ -1,5 +1,6 @@
import logging
from collections.abc import Sequence
from typing import Any
import httpx
from yarl import URL
@@ -19,7 +20,7 @@ def get_plugin_pkg_url(plugin_unique_identifier: str) -> str:
return str((marketplace_api_url / "api/v1/plugins/download").with_query(unique_identifier=plugin_unique_identifier))
def download_plugin_pkg(plugin_unique_identifier: str):
def download_plugin_pkg(plugin_unique_identifier: str) -> bytes:
return download_with_size_limit(get_plugin_pkg_url(plugin_unique_identifier), dify_config.PLUGIN_MAX_PACKAGE_SIZE)
@@ -39,7 +40,7 @@ def batch_fetch_plugin_manifests(plugin_ids: list[str]) -> Sequence[MarketplaceP
return [MarketplacePluginDeclaration.model_validate(plugin) for plugin in response.json()["data"]["plugins"]]
def batch_fetch_plugin_by_ids(plugin_ids: list[str]) -> list[dict]:
def batch_fetch_plugin_by_ids(plugin_ids: list[str]) -> list[dict[str, Any]]:
if not plugin_ids:
return []
@@ -53,10 +54,19 @@ def batch_fetch_plugin_by_ids(plugin_ids: list[str]) -> list[dict]:
response.raise_for_status()
data = response.json()
return data.get("data", {}).get("plugins", [])
plugins = data.get("data", {}).get("plugins", [])
if not isinstance(plugins, list):
raise ValueError("Marketplace did not return a valid plugins list")
result: list[dict[str, Any]] = []
for plugin in plugins:
if not isinstance(plugin, dict) or not all(isinstance(key, str) for key in plugin):
raise ValueError("Marketplace did not return a valid plugins list")
result.append(plugin)
return result
def record_install_plugin_event(plugin_unique_identifier: str):
def record_install_plugin_event(plugin_unique_identifier: str) -> None:
url = str(marketplace_api_url / "api/v1/stats/plugins/install_count")
response = httpx.post(url, json={"unique_identifier": plugin_unique_identifier}, timeout=MARKETPLACE_TIMEOUT)
response.raise_for_status()
+2 -2
View File
@@ -34,7 +34,7 @@ class ProviderCredentialsCache:
else:
return None
def set(self, credentials: dict[str, Any]):
def set(self, credentials: dict[str, Any]) -> None:
"""
Cache model provider credentials.
@@ -43,7 +43,7 @@ class ProviderCredentialsCache:
"""
redis_client.setex(self.cache_key, 86400, json.dumps(credentials))
def delete(self):
def delete(self) -> None:
"""
Delete cached model provider credentials.
+6 -5
View File
@@ -20,17 +20,18 @@ def import_module_from_source[T: (str, bytes)](
raise Exception(f"Failed to load module {module_name} from {py_file_path!r}")
else:
# Refer to: https://docs.python.org/3/library/importlib.html#importing-a-source-file-directly
# FIXME: mypy does not support the type of spec.loader
spec = importlib.util.spec_from_file_location(module_name, py_file_path) # type: ignore[assignment]
if not spec or not spec.loader:
new_spec = importlib.util.spec_from_file_location(module_name, py_file_path)
if not new_spec or not new_spec.loader:
raise Exception(f"Failed to load module {module_name} from {py_file_path!r}")
if use_lazy_loader:
# Refer to: https://docs.python.org/3/library/importlib.html#implementing-lazy-imports
spec.loader = importlib.util.LazyLoader(spec.loader)
new_spec.loader = importlib.util.LazyLoader(new_spec.loader)
spec = new_spec
module = importlib.util.module_from_spec(spec)
if not existed_spec:
sys.modules[module_name] = module
spec.loader.exec_module(module)
if spec.loader is not None:
spec.loader.exec_module(module)
return module
except Exception as e:
logger.exception("Failed to load module %s from script file '%s'", module_name, repr(py_file_path))
+8 -8
View File
@@ -9,11 +9,11 @@ from extensions.ext_redis import redis_client
class ProviderCredentialsCache(ABC):
"""Base class for provider credentials cache"""
def __init__(self, **kwargs):
def __init__(self, **kwargs: Any) -> None:
self.cache_key = self._generate_cache_key(**kwargs)
@abstractmethod
def _generate_cache_key(self, **kwargs) -> str:
def _generate_cache_key(self, **kwargs: Any) -> str:
"""Generate cache key based on subclass implementation"""
pass
@@ -28,11 +28,11 @@ class ProviderCredentialsCache(ABC):
return None
return None
def set(self, config: dict[str, Any]):
def set(self, config: dict[str, Any]) -> None:
"""Cache provider credentials"""
redis_client.setex(self.cache_key, 86400, json.dumps(config))
def delete(self):
def delete(self) -> None:
"""Delete cached provider credentials"""
redis_client.delete(self.cache_key)
@@ -48,7 +48,7 @@ class SingletonProviderCredentialsCache(ProviderCredentialsCache):
)
@override
def _generate_cache_key(self, **kwargs) -> str:
def _generate_cache_key(self, **kwargs: Any) -> str:
tenant_id = kwargs["tenant_id"]
provider_type = kwargs["provider_type"]
identity_name = kwargs["provider_identity"]
@@ -63,7 +63,7 @@ class ToolProviderCredentialsCache(ProviderCredentialsCache):
super().__init__(tenant_id=tenant_id, provider=provider, credential_id=credential_id)
@override
def _generate_cache_key(self, **kwargs) -> str:
def _generate_cache_key(self, **kwargs: Any) -> str:
tenant_id = kwargs["tenant_id"]
provider = kwargs["provider"]
credential_id = kwargs["credential_id"]
@@ -77,10 +77,10 @@ class NoOpProviderCredentialCache:
"""Get cached provider credentials"""
return None
def set(self, config: dict[str, Any]):
def set(self, config: dict[str, Any]) -> None:
"""Cache provider credentials"""
pass
def delete(self):
def delete(self) -> None:
"""Delete cached provider credentials"""
pass
+3 -1
View File
@@ -125,5 +125,7 @@ class ProviderConfigEncrypter:
return data
def create_provider_encrypter(tenant_id: str, config: list[BasicProviderConfig], cache: ProviderConfigCache):
def create_provider_encrypter(
tenant_id: str, config: list[BasicProviderConfig], cache: ProviderConfigCache
) -> tuple[ProviderConfigEncrypter, ProviderConfigCache]:
return ProviderConfigEncrypter(tenant_id=tenant_id, config=config, provider_config_cache=cache), cache
+2 -2
View File
@@ -37,11 +37,11 @@ class ToolParameterCache:
else:
return None
def set(self, parameters: dict[str, Any]):
def set(self, parameters: dict[str, Any]) -> None:
"""Cache model provider credentials."""
redis_client.setex(self.cache_key, 86400, json.dumps(parameters))
def delete(self):
def delete(self) -> None:
"""
Delete cached model provider credentials.
+1 -1
View File
@@ -61,7 +61,7 @@ def get_external_trace_id(request: Any) -> str | None:
return None
def extract_external_trace_id_from_args(args: Mapping[str, Any]):
def extract_external_trace_id_from_args(args: Mapping[str, Any]) -> dict[str, Any]:
"""
Extract 'external_trace_id' from args.
+8
View File
@@ -2,9 +2,13 @@
from core.logging.context import (
clear_request_context,
get_identity_context,
get_request_id,
get_trace_id,
init_request_context,
request_logging_context,
reset_request_context,
set_identity_context,
)
from core.logging.filters import IdentityContextFilter, TraceContextFilter
from core.logging.structured_formatter import StructuredJSONFormatter
@@ -14,7 +18,11 @@ __all__ = [
"StructuredJSONFormatter",
"TraceContextFilter",
"clear_request_context",
"get_identity_context",
"get_request_id",
"get_trace_id",
"init_request_context",
"request_logging_context",
"reset_request_context",
"set_identity_context",
]
+46 -6
View File
@@ -5,10 +5,16 @@ using Python's contextvars for thread-safe and async-safe storage.
"""
import uuid
from contextvars import ContextVar
from collections.abc import Generator
from contextlib import contextmanager
from contextvars import ContextVar, Token
type IdentityContext = tuple[str, str, str]
type LoggingContextTokens = tuple[Token[str], Token[str], Token[IdentityContext]]
_request_id: ContextVar[str] = ContextVar("log_request_id", default="")
_trace_id: ContextVar[str] = ContextVar("log_trace_id", default="")
_identity: ContextVar[IdentityContext] = ContextVar("log_identity", default=("", "", ""))
def get_request_id() -> str:
@@ -21,15 +27,49 @@ def get_trace_id() -> str:
return _trace_id.get()
def init_request_context() -> None:
"""Initialize request context. Call at start of each request."""
def get_identity_context() -> IdentityContext:
"""Get the immutable tenant, user, and user-type snapshot for logging."""
return _identity.get()
def set_identity_context(
*, tenant_id: str | None = None, user_id: str | None = None, user_type: str | None = None
) -> None:
"""Set primitive identity values already resolved by an authentication boundary."""
_identity.set((tenant_id or "", user_id or "", user_type or ""))
def init_request_context() -> LoggingContextTokens:
"""Initialize request context and discard identity left by earlier work."""
req_id = uuid.uuid4().hex[:10]
trace_id = uuid.uuid5(uuid.NAMESPACE_DNS, req_id).hex
_request_id.set(req_id)
_trace_id.set(trace_id)
return (
_request_id.set(req_id),
_trace_id.set(trace_id),
_identity.set(("", "", "")),
)
def reset_request_context(tokens: LoggingContextTokens) -> None:
"""Restore the logging context that preceded ``init_request_context``."""
request_id_token, trace_id_token, identity_token = tokens
_identity.reset(identity_token)
_trace_id.reset(trace_id_token)
_request_id.reset(request_id_token)
@contextmanager
def request_logging_context() -> Generator[None, None, None]:
"""Run work in an isolated logging context and restore its caller afterward."""
tokens = init_request_context()
try:
yield
finally:
reset_request_context(tokens)
def clear_request_context() -> None:
"""Clear request context. Call at end of request (optional)."""
"""Clear request context at a request or task lifecycle boundary."""
_request_id.set("")
_trace_id.set("")
_identity.set(("", "", ""))
+6 -45
View File
@@ -4,10 +4,7 @@ import contextlib
import logging
from typing import override
import flask
from core.logging.context import get_request_id, get_trace_id
from core.logging.structured_formatter import IdentityDict
from core.logging.context import get_identity_context, get_request_id, get_trace_id
class TraceContextFilter(logging.Filter):
@@ -51,49 +48,13 @@ class TraceContextFilter(logging.Filter):
class IdentityContextFilter(logging.Filter):
"""
Filter that adds user identity context to log records.
Extracts tenant_id, user_id, and user_type from Flask-Login current_user.
"""Add an identity snapshot without invoking authentication or database work.
Logging can run while other libraries hold internal locks, so this filter must
only read primitive ContextVar values populated by authentication boundaries.
"""
@override
def filter(self, record: logging.LogRecord) -> bool:
identity = self._extract_identity()
record.tenant_id = identity.get("tenant_id", "")
record.user_id = identity.get("user_id", "")
record.user_type = identity.get("user_type", "")
record.tenant_id, record.user_id, record.user_type = get_identity_context()
return True
def _extract_identity(self) -> IdentityDict:
"""Extract identity from current_user if in request context."""
try:
if not flask.has_request_context():
return {}
from flask_login import current_user
# Check if user is authenticated using the proxy
if not current_user.is_authenticated:
return {}
# Access the underlying user object
user = current_user
from models import Account
from models.model import EndUser
identity: IdentityDict = {}
match user:
case Account():
if user.current_tenant_id:
identity["tenant_id"] = user.current_tenant_id
identity["user_id"] = user.id
identity["user_type"] = "account"
case EndUser():
identity["tenant_id"] = user.tenant_id
identity["user_id"] = user.id
identity["user_type"] = user.type or "end_user"
return identity
except Exception:
return {}
+12 -1
View File
@@ -76,7 +76,11 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
if not user_id:
user = EndUserService.get_or_create_end_user(app)
else:
user = cls._get_user(user_id, app)
try:
user = cls._get_user(user_id, app)
except ValueError:
# Plugins such as WeCom Bot pass external sender IDs rather than EndUser UUIDs.
user = EndUserService.get_or_create_end_user(app, user_id=user_id)
conversation_id = conversation_id or ""
@@ -226,6 +230,13 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
EndUser.app_id == app.id,
)
user = session.scalar(stmt)
if not user:
stmt = select(EndUser).where(
EndUser.session_id == user_id,
EndUser.tenant_id == app.tenant_id,
EndUser.app_id == app.id,
)
user = session.scalar(stmt)
if not user:
stmt = select(Account).where(
Account.id == user_id,
+1 -1
View File
@@ -228,7 +228,7 @@ class CredentialType(enum.StrEnum):
OAUTH2 = "oauth2"
UNAUTHORIZED = "unauthorized"
def get_name(self):
def get_name(self) -> str:
if self == CredentialType.API_KEY:
return "API KEY"
elif self == CredentialType.OAUTH2:
+325 -80
View File
@@ -4,6 +4,8 @@ This module owns plugin daemon management calls that are shared by API services
and core runtimes. Plugin model provider discovery is cached here, alongside
plugin install, uninstall, and upgrade invalidation, so all cache mutations for
plugin-owned provider metadata stay tenant-scoped and in one place.
Provider cache payloads may be stored as prefixed zstd bytes; readers also
accept legacy plain JSON payloads for rolling upgrades and existing Redis keys.
The console plugin list also normalizes endpoint setup counters against live
endpoint records. Some plugin daemon builds return stale ``endpoints_*``
@@ -14,12 +16,15 @@ metadata.
import logging
import time
from collections.abc import Mapping, Sequence
from collections.abc import Iterator, Mapping, Sequence
from contextlib import contextmanager
from mimetypes import guess_type
from typing import ClassVar
from typing import Literal, Protocol, cast
import zstandard
from pydantic import BaseModel, TypeAdapter, ValidationError
from redis import RedisError
from redis.exceptions import LockError
from sqlalchemy import delete, select, update
from sqlalchemy.orm import Session
from yarl import URL
@@ -67,14 +72,18 @@ logger = logging.getLogger(__name__)
_provider_entities_adapter: TypeAdapter[list[ProviderEntity]] = TypeAdapter(list[ProviderEntity])
class PluginService:
_plugin_model_providers_memory_cache: ClassVar[dict[str, tuple[int, float, tuple[ProviderEntity, ...]]]] = {}
class _RedisLock(Protocol):
def acquire(self, *, blocking: bool = True, blocking_timeout: float | None = None) -> bool: ...
def release(self) -> None: ...
class PluginService:
class LatestPluginCache(BaseModel):
plugin_id: str
version: str
unique_identifier: str
status: str
status: Literal["active", "deleted"]
deprecated_reason: str
alternative_plugin_id: str
@@ -82,6 +91,13 @@ class PluginService:
REDIS_TTL = 60 * 5 # 5 minutes
PLUGIN_MODEL_PROVIDERS_REDIS_KEY_PREFIX = "plugin_model_providers:tenant_id:"
PLUGIN_MODEL_PROVIDERS_GENERATION_REDIS_KEY_PREFIX = "plugin_model_providers_generation:tenant_id:"
PLUGIN_MODEL_PROVIDERS_LOCK_REDIS_KEY_PREFIX = "plugin_model_providers_refresh_lock:tenant_id:"
PLUGIN_MODEL_PROVIDERS_REMOTE_DEBUG_REDIS_KEY_PREFIX = "plugin_model_providers_remote_debug:tenant_id:"
PLUGIN_MODEL_PROVIDERS_LOCK_TTL = 30
PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT = 2.0
PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL = 0.05
PLUGIN_MODEL_PROVIDERS_CACHE_COMPRESSION_PREFIX = b"\x00dify-plugin-model-providers-zstd-v1:"
PLUGIN_MODEL_PROVIDERS_CACHE_COMPRESSION_MIN_BYTES = 64 * 1024
PLUGIN_INSTALL_TASK_TERMINAL_STATUSES = (PluginInstallTaskStatus.Success, PluginInstallTaskStatus.Failed)
# Mirror the detail-panel endpoint query size so list reconciliation and
# the visible endpoint drawer exercise the same daemon pagination path.
@@ -98,6 +114,14 @@ class PluginService:
def _get_plugin_model_providers_generation_cache_key(cls, tenant_id: str) -> str:
return f"{cls.PLUGIN_MODEL_PROVIDERS_GENERATION_REDIS_KEY_PREFIX}{tenant_id}"
@classmethod
def _get_plugin_model_providers_lock_key(cls, tenant_id: str, generation: int) -> str:
return f"{cls.PLUGIN_MODEL_PROVIDERS_LOCK_REDIS_KEY_PREFIX}{tenant_id}:generation:{generation}"
@classmethod
def _get_plugin_model_providers_remote_debug_cache_key(cls, tenant_id: str) -> str:
return f"{cls.PLUGIN_MODEL_PROVIDERS_REMOTE_DEBUG_REDIS_KEY_PREFIX}{tenant_id}"
@staticmethod
def _get_provider_short_name_alias(provider: PluginModelProviderEntity) -> str:
"""
@@ -129,8 +153,25 @@ class PluginService:
return declaration
@classmethod
def _copy_provider_entities(cls, providers: Sequence[ProviderEntity]) -> tuple[ProviderEntity, ...]:
return tuple(provider.model_copy(deep=True) for provider in providers)
def _encode_plugin_model_providers_cache_payload(cls, payload: bytes) -> bytes:
if len(payload) < cls.PLUGIN_MODEL_PROVIDERS_CACHE_COMPRESSION_MIN_BYTES:
return payload
return cls.PLUGIN_MODEL_PROVIDERS_CACHE_COMPRESSION_PREFIX + zstandard.compress(payload, level=1)
@classmethod
def _decode_plugin_model_providers_cache_payload(cls, payload: bytes | bytearray | str) -> bytes | bytearray | str:
if isinstance(payload, str):
return payload
prefix = cls.PLUGIN_MODEL_PROVIDERS_CACHE_COMPRESSION_PREFIX
if not payload.startswith(prefix):
return payload
try:
return zstandard.decompress(payload[len(prefix) :])
except zstandard.ZstdError as exc:
raise ValueError("Invalid compressed plugin model providers cache payload.") from exc
@classmethod
def _load_plugin_model_providers_generation(cls, tenant_id: str) -> int | None:
@@ -163,76 +204,35 @@ class PluginService:
return None
@classmethod
def _load_in_memory_plugin_model_providers(
cls, memory_cache_key: str, generation: int
) -> tuple[ProviderEntity, ...] | None:
cached_entry = cls._plugin_model_providers_memory_cache.get(memory_cache_key)
if cached_entry is None:
return None
def _load_cached_plugin_model_providers_for_generation(
cls, tenant_id: str, generation: int | None
) -> tuple[tuple[ProviderEntity, ...] | None, bool]:
if generation is None:
return None, False
cached_generation, expires_at, providers = cached_entry
if cached_generation != generation or time.monotonic() >= expires_at:
cls._plugin_model_providers_memory_cache.pop(memory_cache_key, None)
return None
return cls._copy_provider_entities(providers)
@classmethod
def _store_in_memory_plugin_model_providers(
cls, memory_cache_key: str, generation: int, providers: Sequence[ProviderEntity]
) -> None:
ttl = dify_config.PLUGIN_MODEL_PROVIDERS_CACHE_TTL
if ttl <= 0:
cls._plugin_model_providers_memory_cache.pop(memory_cache_key, None)
return
cls._plugin_model_providers_memory_cache[memory_cache_key] = (
generation,
time.monotonic() + ttl,
cls._copy_provider_entities(providers),
)
@classmethod
def _load_cached_plugin_model_providers(
cls, tenant_id: str, *, client: PluginModelClient | None = None
) -> tuple[ProviderEntity, ...] | None:
generation = cls._load_plugin_model_providers_generation(tenant_id)
if generation is not None:
in_memory_cached_providers = cls._load_in_memory_plugin_model_providers(tenant_id, generation)
if in_memory_cached_providers is not None:
return in_memory_cached_providers
cache_keys = []
if generation is not None:
cache_keys.append(cls._get_plugin_model_providers_cache_key(tenant_id, generation))
if generation == 0:
cache_keys.append(cls._get_plugin_model_providers_cache_key(tenant_id))
if not cache_keys:
return None
cache_keys = [cls._get_plugin_model_providers_cache_key(tenant_id, generation)]
try:
cached_provider_entries = redis_client.mget(cache_keys)
except (RedisError, RuntimeError):
except (LockError, RedisError, RuntimeError):
logger.warning("Failed to read cached plugin model providers for tenant %s.", tenant_id, exc_info=True)
return None
return None, False
if len(cached_provider_entries) != len(cache_keys):
logger.warning(
"Unexpected cached plugin model providers response size for tenant %s.",
tenant_id,
)
return None
return None, False
for cache_key, cached_providers in zip(cache_keys, cached_provider_entries):
if not cached_providers:
continue
try:
providers = tuple(_provider_entities_adapter.validate_json(cached_providers))
if generation is not None:
cls._store_in_memory_plugin_model_providers(tenant_id, generation, providers)
return providers
payload = cls._decode_plugin_model_providers_cache_payload(cached_providers)
providers = tuple(_provider_entities_adapter.validate_json(payload))
return providers, True
except (TypeError, ValueError, ValidationError):
logger.warning(
"Invalid cached plugin model providers for tenant %s; deleting cache key %s.",
@@ -249,7 +249,7 @@ class PluginService:
exc_info=True,
)
return None
return None, True
@classmethod
def _store_cached_plugin_model_providers(
@@ -257,15 +257,199 @@ class PluginService:
) -> None:
cache_key = cls._get_plugin_model_providers_cache_key(tenant_id, generation)
try:
payload = _provider_entities_adapter.dump_json(list(providers)).decode("utf-8")
payload = cls._encode_plugin_model_providers_cache_payload(
_provider_entities_adapter.dump_json(list(providers))
)
redis_client.setex(cache_key, dify_config.PLUGIN_MODEL_PROVIDERS_CACHE_TTL, payload)
except (RedisError, RuntimeError):
logger.warning("Failed to cache plugin model providers for tenant %s.", tenant_id, exc_info=True)
@classmethod
def _get_remote_model_plugin_cache_marker(cls, plugins: Sequence[PluginEntity]) -> str | None:
remote_model_plugins = sorted(
f"{plugin.plugin_id}:{plugin.plugin_unique_identifier}"
for plugin in plugins
if plugin.source == PluginInstallationSource.Remote
)
if not remote_model_plugins:
return None
return "\n".join(remote_model_plugins)
@classmethod
def _load_cached_remote_model_plugin_marker(cls, tenant_id: str) -> str | None:
cache_key = cls._get_plugin_model_providers_remote_debug_cache_key(tenant_id)
try:
cached_marker = redis_client.get(cache_key)
except (RedisError, RuntimeError):
logger.warning("Failed to read remote debug model plugin marker for tenant %s.", tenant_id, exc_info=True)
return None
if cached_marker is None:
return None
if isinstance(cached_marker, bytes):
try:
return cached_marker.decode()
except UnicodeDecodeError:
logger.warning(
"Invalid remote debug model plugin marker for tenant %s; deleting cache marker.",
tenant_id,
exc_info=True,
)
try:
redis_client.delete(cache_key)
except (RedisError, RuntimeError):
logger.warning(
"Failed to delete invalid remote debug model plugin marker for tenant %s.",
tenant_id,
exc_info=True,
)
return None
if isinstance(cached_marker, str):
return cached_marker
logger.warning("Invalid remote debug model plugin marker for tenant %s; deleting cache marker.", tenant_id)
try:
redis_client.delete(cache_key)
except (RedisError, RuntimeError):
logger.warning(
"Failed to delete invalid remote debug model plugin marker for tenant %s.",
tenant_id,
exc_info=True,
)
return None
@classmethod
def _store_cached_remote_model_plugin_marker(cls, tenant_id: str, marker: str | None) -> None:
cache_key = cls._get_plugin_model_providers_remote_debug_cache_key(tenant_id)
try:
if marker is None:
redis_client.delete(cache_key)
else:
redis_client.setex(cache_key, dify_config.PLUGIN_MODEL_PROVIDERS_CACHE_TTL, marker)
except (RedisError, RuntimeError):
logger.warning("Failed to cache remote debug model plugin marker for tenant %s.", tenant_id, exc_info=True)
@classmethod
def _load_cached_plugin_model_provider_plugin_ids(cls, tenant_id: str) -> set[str] | None:
"""Return plugin ids represented by the current provider cache, or None when no usable cache exists."""
generation = cls._load_plugin_model_providers_generation(tenant_id)
cached_providers, _ = cls._load_cached_plugin_model_providers_for_generation(tenant_id, generation)
if cached_providers is None:
return None
plugin_ids: set[str] = set()
for provider in cached_providers:
last_slash = provider.provider.rfind("/")
if last_slash > 0:
plugin_ids.add(provider.provider[:last_slash])
return plugin_ids
@classmethod
def _should_invalidate_model_provider_cache_for_remote_model_plugins(
cls,
tenant_id: str,
plugins: Sequence[PluginEntity],
) -> bool:
remote_model_plugin_marker = cls._get_remote_model_plugin_cache_marker(plugins)
cached_remote_model_plugin_marker = cls._load_cached_remote_model_plugin_marker(tenant_id)
if remote_model_plugin_marker is None:
return cached_remote_model_plugin_marker is not None
if remote_model_plugin_marker != cached_remote_model_plugin_marker:
return True
remote_model_plugin_ids = {
plugin.plugin_id for plugin in plugins if plugin.source == PluginInstallationSource.Remote
}
cached_plugin_ids = cls._load_cached_plugin_model_provider_plugin_ids(tenant_id)
if cached_plugin_ids is None:
return False
return not remote_model_plugin_ids.issubset(cached_plugin_ids)
@classmethod
@contextmanager
def _plugin_model_providers_refresh_lock(
cls, tenant_id: str, generation: int, *, wait_timeout: float
) -> Iterator[bool]:
lock_key = cls._get_plugin_model_providers_lock_key(tenant_id, generation)
try:
refresh_lock: _RedisLock = redis_client.lock(
lock_key,
timeout=cls.PLUGIN_MODEL_PROVIDERS_LOCK_TTL,
sleep=cls.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL,
)
except (RedisError, RuntimeError):
logger.warning(
"Failed to create plugin model providers refresh lock for tenant %s.",
tenant_id,
exc_info=True,
)
yield False
return
try:
lock_acquired = refresh_lock.acquire(blocking=True, blocking_timeout=wait_timeout)
except LockError:
logger.warning(
"Provider refresh lock timed out; direct daemon fallback. tenant_id=%s generation=%s",
tenant_id,
generation,
exc_info=True,
)
yield False
return
except (RedisError, RuntimeError):
# Redis failures should not block provider discovery; callers fetch directly from the daemon.
logger.warning(
"Failed to acquire plugin model providers refresh lock for tenant %s.",
tenant_id,
exc_info=True,
)
yield False
return
if not lock_acquired:
logger.warning(
"Provider refresh lock timed out; direct daemon fallback. tenant_id=%s generation=%s",
tenant_id,
generation,
)
yield False
return
try:
yield True
finally:
try:
refresh_lock.release()
except (LockError, RedisError, RuntimeError):
# Release failures must not hide the daemon result or the original exception.
logger.warning(
"Failed to release plugin model providers refresh lock for tenant %s generation %s.",
tenant_id,
generation,
exc_info=True,
)
@classmethod
def _fetch_and_cache_plugin_model_providers(
cls, tenant_id: str, client: PluginModelClient | None, *, refresh_generation: int | None
) -> tuple[ProviderEntity, ...]:
model_client = client or PluginModelClient()
providers = tuple(
cls._to_provider_entity(provider) for provider in model_client.fetch_model_providers(tenant_id)
)
generation = cls._load_plugin_model_providers_generation(tenant_id)
if generation is not None and generation == refresh_generation:
cls._store_cached_plugin_model_providers(tenant_id, generation, providers)
return providers
@classmethod
def invalidate_plugin_model_providers_cache(cls, tenant_id: str) -> None:
"""Invalidate tenant-scoped provider metadata across Redis and worker-local mirrors."""
cls._plugin_model_providers_memory_cache.pop(tenant_id, None)
"""Invalidate tenant-scoped provider metadata stored in Redis."""
cache_key = cls._get_plugin_model_providers_cache_key(tenant_id)
generation_key = cls._get_plugin_model_providers_generation_cache_key(tenant_id)
try:
@@ -287,21 +471,68 @@ class PluginService:
are intentionally owned by this service so tenant isolation and cache
expiry are handled in one place.
"""
cached_providers = cls._load_cached_plugin_model_providers(tenant_id, client=client)
if cached_providers is not None:
return cached_providers
deadline = time.monotonic() + cls.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT
model_client = client or PluginModelClient()
providers = tuple(
cls._to_provider_entity(provider) for provider in model_client.fetch_model_providers(tenant_id)
)
if not providers:
return providers
generation = cls._load_plugin_model_providers_generation(tenant_id)
if generation is not None:
cls._store_in_memory_plugin_model_providers(tenant_id, generation, providers)
cls._store_cached_plugin_model_providers(tenant_id, generation, providers)
return providers
while True:
generation = cls._load_plugin_model_providers_generation(tenant_id)
cached_providers, cache_available = cls._load_cached_plugin_model_providers_for_generation(
tenant_id, generation
)
if cached_providers is not None:
return cached_providers
if generation is None or not cache_available:
return cls._fetch_and_cache_plugin_model_providers(
tenant_id,
client,
refresh_generation=generation,
)
wait_timeout = deadline - time.monotonic()
if wait_timeout < 0:
logger.warning(
"Provider refresh lock timed out; direct daemon fallback. tenant_id=%s generation=%s",
tenant_id,
generation,
)
return cls._fetch_and_cache_plugin_model_providers(
tenant_id,
client,
refresh_generation=generation,
)
with cls._plugin_model_providers_refresh_lock(
tenant_id,
generation,
wait_timeout=wait_timeout,
) as lock_acquired:
if not lock_acquired:
return cls._fetch_and_cache_plugin_model_providers(
tenant_id,
client,
refresh_generation=generation,
)
latest_generation = cls._load_plugin_model_providers_generation(tenant_id)
cached_providers, cache_available = cls._load_cached_plugin_model_providers_for_generation(
tenant_id, latest_generation
)
if cached_providers is not None:
return cached_providers
if latest_generation is None or not cache_available:
return cls._fetch_and_cache_plugin_model_providers(
tenant_id,
client,
refresh_generation=latest_generation,
)
if latest_generation != generation:
continue
return cls._fetch_and_cache_plugin_model_providers(
tenant_id,
client,
refresh_generation=generation,
)
@staticmethod
def fetch_latest_plugin_version(plugin_ids: Sequence[str]) -> Mapping[str, LatestPluginCache | None]:
@@ -340,7 +571,7 @@ class PluginService:
plugin_id=plugin_id,
version=manifest.latest_version,
unique_identifier=manifest.latest_package_identifier,
status=manifest.status,
status=cast(Literal["active", "deleted"], manifest.status),
deprecated_reason=manifest.deprecated_reason,
alternative_plugin_id=manifest.alternative_plugin_id,
)
@@ -450,7 +681,21 @@ class PluginService:
This keeps pagination usable before category is persisted on installation rows.
"""
manager = PluginInstaller()
return manager.list_plugins_by_category(tenant_id, category, page, page_size)
plugins = manager.list_plugins_by_category(tenant_id, category, page, page_size)
if category == PluginCategory.Model:
should_invalidate_model_provider_cache = (
PluginService._should_invalidate_model_provider_cache_for_remote_model_plugins(
tenant_id,
plugins.list,
)
)
if should_invalidate_model_provider_cache:
PluginService.invalidate_plugin_model_providers_cache(tenant_id)
remote_model_plugin_marker = PluginService._get_remote_model_plugin_cache_marker(plugins.list)
PluginService._store_cached_remote_model_plugin_marker(tenant_id, remote_model_plugin_marker)
return plugins
@staticmethod
def _normalize_endpoint_count(value: object) -> int:
@@ -44,11 +44,13 @@ def _load_plugin_factory(vector_type: str) -> type[AbstractVectorFactory] | None
for ep in _vector_backend_entry_points():
if ep.name != vector_type:
continue
logger.info("loading vector backend entry point %s (%s)", ep.name, ep.value)
try:
loaded = ep.load()
except Exception:
logger.exception("Failed to load vector backend entry point %s", ep.name)
raise
logger.info("vector backend entry point %s loaded", ep.name)
return loaded # type: ignore[return-value]
return None
@@ -77,6 +79,7 @@ def get_vector_factory_class(vector_type: str) -> type[AbstractVectorFactory]:
if vector_type in _VECTOR_FACTORY_CACHE:
return _VECTOR_FACTORY_CACHE[vector_type]
logger.info("vector backend cache miss, discovering entry points: %s", vector_type)
plugin_cls = _load_plugin_factory(vector_type)
if plugin_cls is not None:
_VECTOR_FACTORY_CACHE[vector_type] = plugin_cls
@@ -144,7 +144,20 @@ class Vector:
if not vector_type:
raise ValueError("Vector store must be specified.")
logger.info(
"resolving vector backend factory: type=%s tenant_id=%s dataset_id=%s",
vector_type,
self._dataset.tenant_id,
self._dataset.id,
)
vector_factory_cls = self.get_vector_factory(vector_type)
logger.info(
"vector backend factory resolved: type=%s cls=%s tenant_id=%s dataset_id=%s",
vector_type,
vector_factory_cls.__name__,
self._dataset.tenant_id,
self._dataset.id,
)
return vector_factory_cls().init_vector(self._dataset, self._attributes, self._embeddings)
@staticmethod
+14 -4
View File
@@ -1030,6 +1030,10 @@ class DatasetRetrieval:
):
"""
Persist dataset query audit rows for retrieval requests.
Query audit logging is a side effect of retrieval. Keep it in an
independent transaction so failures or commits here do not affect the
request/workflow transaction that called the retriever.
"""
if not query and not attachment_ids:
return
@@ -1041,6 +1045,9 @@ class DatasetRetrieval:
app_id,
)
return
created_by_role = self._resolve_creator_user_role(user_from)
if created_by_role is None:
return
dataset_queries = []
for dataset_id in dataset_ids:
contents = []
@@ -1055,13 +1062,16 @@ class DatasetRetrieval:
content=json.dumps(contents),
source=DatasetQuerySource.APP,
source_app_id=app_id,
created_by_role=CreatorUserRole(user_from),
created_by_role=created_by_role,
created_by=created_by,
)
dataset_queries.append(dataset_query)
if dataset_queries:
db.session.add_all(dataset_queries)
db.session.commit()
if not dataset_queries:
return
with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session:
session.add_all(dataset_queries)
def _retriever(
self,
+2 -2
View File
@@ -1,3 +1,3 @@
from core.rbac.entities import RBACPermission, RBACResourceScope
from core.rbac.entities import RBACPermission, RBACResourceScope, RBACResourceWhitelistScope
__all__ = ["RBACPermission", "RBACResourceScope"]
__all__ = ["RBACPermission", "RBACResourceScope", "RBACResourceWhitelistScope"]
+8
View File
@@ -13,6 +13,14 @@ class RBACResourceScope(StrEnum):
WORKSPACE = "workspace"
class RBACResourceWhitelistScope(StrEnum):
"""Whitelist scopes accepted by RBAC app and dataset access config APIs."""
ALL = "all"
SPECIFIC = "specific"
ONLY_ME = "only_me"
class RBACPermission(StrEnum):
"""Permission points (RBAC scenes) checked by ``rbac_permission_required``.
@@ -13,7 +13,6 @@ from sqlalchemy.orm import sessionmaker
from core.repositories.factory import WorkflowExecutionRepository
from graphon.entities import WorkflowExecution
from libs.helper import extract_tenant_id
from models import Account, CreatorUserRole, EndUser
from models.enums import WorkflowRunTriggeredFrom
from tasks.workflow_execution_tasks import (
@@ -47,6 +46,7 @@ class CeleryWorkflowExecutionRepository(WorkflowExecutionRepository):
def __init__(
self,
session_factory: sessionmaker | Engine,
tenant_id: str,
user: Account | EndUser,
app_id: str | None,
triggered_from: WorkflowRunTriggeredFrom | None,
@@ -56,7 +56,8 @@ class CeleryWorkflowExecutionRepository(WorkflowExecutionRepository):
Args:
session_factory: SQLAlchemy sessionmaker or engine for fallback operations
user: Account or EndUser object containing tenant_id, user ID, and role information
tenant_id: Tenant that owns the workflow execution
user: Account or EndUser used for creator attribution
app_id: App ID for filtering by application (can be None)
triggered_from: Source of the execution trigger (DEBUGGING or APP_RUN)
"""
@@ -71,10 +72,8 @@ class CeleryWorkflowExecutionRepository(WorkflowExecutionRepository):
f"Invalid session_factory type {type(session_factory).__name__}; expected sessionmaker or Engine"
)
# Extract tenant_id from user
tenant_id = extract_tenant_id(user)
if not tenant_id:
raise ValueError("User must have a tenant_id or current_tenant_id")
raise ValueError("tenant_id is required")
self._tenant_id = tenant_id
# Store app context
@@ -17,7 +17,6 @@ from core.repositories.factory import (
WorkflowNodeExecutionRepository,
)
from graphon.entities import WorkflowNodeExecution
from libs.helper import extract_tenant_id
from models import Account, CreatorUserRole, EndUser
from models.workflow import WorkflowNodeExecutionTriggeredFrom
from tasks.workflow_node_execution_tasks import (
@@ -54,6 +53,7 @@ class CeleryWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository):
def __init__(
self,
session_factory: sessionmaker | Engine,
tenant_id: str,
user: Account | EndUser,
app_id: str | None,
triggered_from: WorkflowNodeExecutionTriggeredFrom | None,
@@ -63,7 +63,8 @@ class CeleryWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository):
Args:
session_factory: SQLAlchemy sessionmaker or engine for fallback operations
user: Account or EndUser object containing tenant_id, user ID, and role information
tenant_id: Tenant that owns the workflow node execution
user: Account or EndUser used for creator attribution
app_id: App ID for filtering by application (can be None)
triggered_from: Source of the execution trigger (SINGLE_STEP or WORKFLOW_RUN)
"""
@@ -78,10 +79,8 @@ class CeleryWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository):
f"Invalid session_factory type {type(session_factory).__name__}; expected sessionmaker or Engine"
)
# Extract tenant_id from user
tenant_id = extract_tenant_id(user)
if not tenant_id:
raise ValueError("User must have a tenant_id or current_tenant_id")
raise ValueError("tenant_id is required")
self._tenant_id = tenant_id
# Store app context
+8 -2
View File
@@ -62,6 +62,7 @@ class DifyCoreRepositoryFactory:
def create_workflow_execution_repository(
cls,
session_factory: sessionmaker | Engine,
tenant_id: str,
user: Account | EndUser,
app_id: str,
triggered_from: WorkflowRunTriggeredFrom,
@@ -71,7 +72,8 @@ class DifyCoreRepositoryFactory:
Args:
session_factory: SQLAlchemy sessionmaker or engine
user: Account or EndUser object
tenant_id: Tenant that owns the workflow execution
user: Account or EndUser used for creator attribution
app_id: Application ID
triggered_from: Source of the execution trigger
@@ -87,6 +89,7 @@ class DifyCoreRepositoryFactory:
repository_class = import_string(class_path)
return repository_class(
session_factory=session_factory,
tenant_id=tenant_id,
user=user,
app_id=app_id,
triggered_from=triggered_from,
@@ -98,6 +101,7 @@ class DifyCoreRepositoryFactory:
def create_workflow_node_execution_repository(
cls,
session_factory: sessionmaker | Engine,
tenant_id: str,
user: Account | EndUser,
app_id: str,
triggered_from: WorkflowNodeExecutionTriggeredFrom,
@@ -107,7 +111,8 @@ class DifyCoreRepositoryFactory:
Args:
session_factory: SQLAlchemy sessionmaker or engine
user: Account or EndUser object
tenant_id: Tenant that owns the workflow node execution
user: Account or EndUser used for creator attribution
app_id: Application ID
triggered_from: Source of the execution trigger
@@ -123,6 +128,7 @@ class DifyCoreRepositoryFactory:
repository_class = import_string(class_path)
return repository_class(
session_factory=session_factory,
tenant_id=tenant_id,
user=user,
app_id=app_id,
triggered_from=triggered_from,
@@ -13,7 +13,6 @@ from core.repositories.factory import WorkflowExecutionRepository
from graphon.entities import WorkflowExecution
from graphon.enums import WorkflowExecutionStatus, WorkflowType
from graphon.workflow_type_encoder import WorkflowRuntimeTypeConverter
from libs.helper import extract_tenant_id
from models import (
Account,
CreatorUserRole,
@@ -40,6 +39,7 @@ class SQLAlchemyWorkflowExecutionRepository(WorkflowExecutionRepository):
def __init__(
self,
session_factory: sessionmaker | Engine,
tenant_id: str,
user: Account | EndUser,
app_id: str | None,
triggered_from: WorkflowRunTriggeredFrom | None,
@@ -49,7 +49,8 @@ class SQLAlchemyWorkflowExecutionRepository(WorkflowExecutionRepository):
Args:
session_factory: SQLAlchemy sessionmaker or engine for creating sessions
user: Account or EndUser object containing tenant_id, user ID, and role information
tenant_id: Tenant that owns the workflow execution
user: Account or EndUser used for creator attribution
app_id: App ID for filtering by application (can be None)
triggered_from: Source of the execution trigger (DEBUGGING or APP_RUN)
"""
@@ -64,10 +65,8 @@ class SQLAlchemyWorkflowExecutionRepository(WorkflowExecutionRepository):
f"Invalid session_factory type {type(session_factory).__name__}; expected sessionmaker or Engine"
)
# Extract tenant_id from user
tenant_id = extract_tenant_id(user)
if not tenant_id:
raise ValueError("User must have a tenant_id or current_tenant_id")
raise ValueError("tenant_id is required")
self._tenant_id = tenant_id
# Store app context
@@ -23,7 +23,6 @@ from graphon.entities import WorkflowNodeExecution
from graphon.enums import WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus
from graphon.model_runtime.utils.encoders import jsonable_encoder
from graphon.workflow_type_encoder import WorkflowRuntimeTypeConverter
from libs.helper import extract_tenant_id
from libs.uuid_utils import uuidv7
from models import (
Account,
@@ -63,6 +62,7 @@ class SQLAlchemyWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository)
def __init__(
self,
session_factory: sessionmaker | Engine,
tenant_id: str,
user: Account | EndUser,
app_id: str | None,
triggered_from: WorkflowNodeExecutionTriggeredFrom | None,
@@ -72,7 +72,8 @@ class SQLAlchemyWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository)
Args:
session_factory: SQLAlchemy sessionmaker or engine for creating sessions
user: Account or EndUser object containing tenant_id, user ID, and role information
tenant_id: Tenant that owns the workflow node execution
user: Account or EndUser used for creator attribution
app_id: App ID for filtering by application (can be None)
triggered_from: Source of the execution trigger (SINGLE_STEP or WORKFLOW_RUN)
"""
@@ -87,10 +88,8 @@ class SQLAlchemyWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository)
f"Invalid session_factory type {type(session_factory).__name__}; expected sessionmaker or Engine"
)
# Extract tenant_id from user
tenant_id = extract_tenant_id(user)
if not tenant_id:
raise ValueError("User must have a tenant_id or current_tenant_id")
raise ValueError("tenant_id is required")
self._tenant_id = tenant_id
# Store app context
@@ -299,6 +298,7 @@ class SQLAlchemyWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository)
content=value_json.encode("utf-8"),
mimetype="application/json",
user=self._user,
tenant_id=self._tenant_id,
)
offload = WorkflowNodeExecutionOffload(
id=uuidv7(),
+6 -5
View File
@@ -102,16 +102,17 @@ class ApiTool(Tool):
elif not isinstance(credentials["api_key_value"], str):
raise ToolProviderCredentialValidationError("api_key_value must be a string")
api_key_value = credentials["api_key_value"]
if "api_key_header_prefix" in credentials:
api_key_header_prefix = credentials["api_key_header_prefix"]
if api_key_header_prefix == "basic" and credentials["api_key_value"]:
credentials["api_key_value"] = f"Basic {credentials['api_key_value']}"
elif api_key_header_prefix == "bearer" and credentials["api_key_value"]:
credentials["api_key_value"] = f"Bearer {credentials['api_key_value']}"
if api_key_header_prefix == "basic" and api_key_value:
api_key_value = f"Basic {api_key_value}"
elif api_key_header_prefix == "bearer" and api_key_value:
api_key_value = f"Bearer {api_key_value}"
elif api_key_header_prefix == "custom":
pass
headers[api_key_header] = credentials["api_key_value"]
headers[api_key_header] = api_key_value
elif credentials["auth_type"] == "api_key_query":
# For query parameter authentication, we don't add anything to headers
+39 -36
View File
@@ -7,6 +7,7 @@ from datetime import UTC, datetime
from mimetypes import guess_type
from typing import Any, Union, cast
from sqlalchemy.orm import sessionmaker
from yarl import URL
from core.app.entities.app_invoke_entities import InvokeFrom
@@ -338,47 +339,49 @@ class ToolEngine:
user_id: str,
) -> list[str]:
"""
Create message file
Create message files produced by a tool call.
Tool file persistence is a side effect of agent execution. Use an
independent transaction so this helper never commits or closes the
caller's request-scoped session.
:return: message file ids
"""
result = []
for message in tool_messages:
if "image" in message.mimetype:
file_type = FileType.IMAGE
elif "video" in message.mimetype:
file_type = FileType.VIDEO
elif "audio" in message.mimetype:
file_type = FileType.AUDIO
elif "text" in message.mimetype or "pdf" in message.mimetype:
file_type = FileType.DOCUMENT
else:
file_type = FileType.CUSTOM
with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session:
for message in tool_messages:
# extract tool file id from url
tool_file_id = message.url.split("/")[-1].split(".")[0]
message_file = MessageFile(
message_id=agent_message.id,
type=ToolEngine._resolve_tool_file_type(message),
transfer_method=FileTransferMethod.TOOL_FILE,
belongs_to=MessageFileBelongsTo.ASSISTANT,
url=message.url,
upload_file_id=tool_file_id,
created_by_role=(
CreatorUserRole.ACCOUNT
if invoke_from in {InvokeFrom.EXPLORE, InvokeFrom.DEBUGGER}
else CreatorUserRole.END_USER
),
created_by=user_id,
)
# extract tool file id from url
tool_file_id = message.url.split("/")[-1].split(".")[0]
message_file = MessageFile(
message_id=agent_message.id,
type=file_type,
transfer_method=FileTransferMethod.TOOL_FILE,
belongs_to=MessageFileBelongsTo.ASSISTANT,
url=message.url,
upload_file_id=tool_file_id,
created_by_role=(
CreatorUserRole.ACCOUNT
if invoke_from in {InvokeFrom.EXPLORE, InvokeFrom.DEBUGGER}
else CreatorUserRole.END_USER
),
created_by=user_id,
)
db.session.add(message_file)
db.session.commit()
db.session.refresh(message_file)
result.append(message_file.id)
db.session.close()
session.add(message_file)
result.append(message_file.id)
return result
@staticmethod
def _resolve_tool_file_type(message: ToolInvokeMessageBinary) -> FileType:
if "image" in message.mimetype:
return FileType.IMAGE
elif "video" in message.mimetype:
return FileType.VIDEO
elif "audio" in message.mimetype:
return FileType.AUDIO
elif "text" in message.mimetype or "pdf" in message.mimetype:
return FileType.DOCUMENT
else:
return FileType.CUSTOM
+2 -2
View File
@@ -5,10 +5,10 @@ from json import loads as json_loads
from json.decoder import JSONDecodeError
from typing import Any, TypedDict
import httpx
from flask import has_request_context, request
from yaml import YAMLError, safe_load
from core.helper import ssrf_proxy
from core.tools.entities.common_entities import I18nObject
from core.tools.entities.tool_bundle import ApiToolBundle
from core.tools.entities.tool_entities import ApiProviderSchemaType, ToolParameter
@@ -376,7 +376,7 @@ class ApiBasedToolSchemaParser:
raise ToolNotSupportedError("Only openapi is supported now.")
# get openapi yaml
response = httpx.get(
response = ssrf_proxy.get(
api_url, headers={"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) "}, timeout=5
)
+15 -3
View File
@@ -46,18 +46,26 @@ logger = logging.getLogger(__name__)
_file_access_controller = DatabaseFileAccessController()
def iter_dify_graph_engine_events(engine: GraphEngine) -> Generator[GraphEngineEvent, None, None]:
def iter_dify_graph_engine_events(
engine: GraphEngine,
response_stream_filter: ResponseStreamFilter | None = None,
) -> Generator[GraphEngineEvent, None, None]:
"""
Apply Dify's response streaming compatibility filter to GraphEngine events.
Graphon v0.5.0 emits raw variable stream chunks and requires callers to opt
into the legacy response-ordered stream behavior that Dify exposes to its
workflow runners and tests.
``response_stream_filter``, when supplied, must be the same instance a
caller intends to persist on pause (see ``PauseStatePersistenceLayer``) so
the filter's ``paths_map`` reflects everything the engine has actually
streamed for this run.
"""
yield from filter_graph_events(
engine.run(),
context=GraphEventFilterContext.from_engine(engine),
filters=[ResponseStreamFilter()],
filters=[response_stream_filter or ResponseStreamFilter()],
)
@@ -167,6 +175,7 @@ class WorkflowEntry:
variable_pool: VariablePool,
graph_runtime_state: GraphRuntimeState,
command_channel: CommandChannel | None = None,
response_stream_filter: ResponseStreamFilter | None = None,
) -> None:
"""
Init workflow entry
@@ -183,6 +192,8 @@ class WorkflowEntry:
:param variable_pool: variable pool
:param graph_runtime_state: pre-created graph runtime state
:param command_channel: command channel for external control (optional, defaults to InMemoryChannel)
:param response_stream_filter: pre-restored filter for resumed runs (optional, defaults to a fresh
ResponseStreamFilter for runs with no prior pause)
:param thread_pool_id: thread pool id
"""
# check call depth
@@ -195,6 +206,7 @@ class WorkflowEntry:
command_channel = InMemoryChannel()
self.command_channel = command_channel
self._response_stream_filter = response_stream_filter or ResponseStreamFilter()
execution_context = capture_current_context()
graph_runtime_state.execution_context = execution_context
self._child_engine_builder = _WorkflowChildEngineBuilder(tenant_id=tenant_id)
@@ -240,7 +252,7 @@ class WorkflowEntry:
try:
# Preserve Dify's response-stream semantics on top of Graphon 0.5.0.
generator = iter_dify_graph_engine_events(graph_engine)
generator = iter_dify_graph_engine_events(graph_engine, self._response_stream_filter)
yield from generator
except GenerateTaskStoppedError:
pass
+2 -2
View File
@@ -35,10 +35,10 @@ if [[ "${MODE}" == "worker" ]]; then
if [[ -z "${CELERY_QUEUES}" ]]; then
if [[ "${EDITION}" == "CLOUD" ]]; then
# Cloud edition: separate queues for dataset and trigger tasks
DEFAULT_QUEUES="api_token,dataset,dataset_summary,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,plugin,workflow_storage,conversation,workflow_professional,workflow_team,workflow_sandbox,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_publisher,trigger_refresh_executor,retention,workflow_based_app_execution"
DEFAULT_QUEUES="api_token,dataset,dataset_summary,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,app_rbac,plugin,workflow_storage,conversation,workflow_professional,workflow_team,workflow_sandbox,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_publisher,trigger_refresh_executor,retention,workflow_based_app_execution"
else
# Community edition (SELF_HOSTED): dataset, pipeline and workflow have separate queues
DEFAULT_QUEUES="api_token,dataset,dataset_summary,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,plugin,workflow_storage,conversation,workflow,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_publisher,trigger_refresh_executor,retention,workflow_based_app_execution"
DEFAULT_QUEUES="api_token,dataset,dataset_summary,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,app_rbac,plugin,workflow_storage,conversation,workflow,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_publisher,trigger_refresh_executor,retention,workflow_based_app_execution"
fi
else
DEFAULT_QUEUES="${CELERY_QUEUES}"
+6 -5
View File
@@ -98,12 +98,12 @@ def get_celery_redis_global_keyprefix() -> str | None:
def init_app(app: DifyApp) -> Celery:
class FlaskTask(Task):
def __call__(self, *args: object, **kwargs: object) -> object:
from core.logging.context import init_request_context
from core.logging.context import request_logging_context
with app.app_context():
# Initialize logging context for this task (similar to before_request in Flask)
init_request_context()
return self.run(*args, **kwargs)
# Keep task identity through Flask teardown logging and restore callers in eager/nested execution.
with request_logging_context():
with app.app_context():
return self.run(*args, **kwargs)
broker_transport_options = get_celery_broker_transport_options()
@@ -153,6 +153,7 @@ def init_app(app: DifyApp) -> Celery:
"tasks.trigger_processing_tasks", # async trigger processing
"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.app_generate.resume_agent_app_task", # ENG-635: Agent v2 chat ask_human resume
]
day = dify_config.CELERY_BEAT_SCHEDULER_TIME
+2
View File
@@ -27,6 +27,7 @@ def init_app(app: DifyApp):
install_plugins,
install_rag_pipeline_plugins,
migrate_data_for_plugin,
migrate_dataset_permissions_to_rbac,
migrate_member_roles_to_rbac,
migrate_oss,
migration_data_wizard,
@@ -56,6 +57,7 @@ def init_app(app: DifyApp):
upgrade_db,
fix_app_site_missing,
migrate_data_for_plugin,
migrate_dataset_permissions_to_rbac,
migrate_member_roles_to_rbac,
backfill_plugin_auto_upgrade,
extract_plugins,
+12 -3
View File
@@ -9,6 +9,7 @@ from werkzeug.exceptions import NotFound, Unauthorized
from configs import dify_config
from constants import HEADER_NAME_APP_CODE
from core.logging.context import set_identity_context
from dify_app import DifyApp
from extensions.ext_database import db
from libs.passport import PassportService
@@ -149,13 +150,21 @@ def load_user_from_request(request_from_flask_login: Request) -> LoginUser | Non
@user_logged_in.connect
@user_loaded_from_request.connect
def on_user_logged_in(_sender: object, user: LoginUser) -> None:
"""Called when a user logged in.
"""Snapshot authenticated identity into the side-effect-free logging context.
Note: AccountService.load_logged_in_account will populate user.current_tenant_id
through the load_user method, which calls account.set_tenant_id().
"""
# tenant_id context variable removed - using current_user.current_tenant_id directly
pass
set_identity_context()
try:
match user:
case Account():
set_identity_context(tenant_id=user.current_tenant_id, user_id=user.id, user_type="account")
case EndUser():
set_identity_context(tenant_id=user.tenant_id, user_id=user.id, user_type=user.type or "end_user")
except Exception:
# Logging enrichment must never make authentication fail.
return
@login_manager.unauthorized_handler
@@ -12,7 +12,6 @@ from core.repositories.sqlalchemy_workflow_execution_repository import SQLAlchem
from extensions.logstore.aliyun_logstore import AliyunLogStore
from graphon.entities import WorkflowExecution
from graphon.workflow_type_encoder import WorkflowRuntimeTypeConverter
from libs.helper import extract_tenant_id
from models import (
Account,
CreatorUserRole,
@@ -27,6 +26,7 @@ class LogstoreWorkflowExecutionRepository(WorkflowExecutionRepository):
def __init__(
self,
session_factory: sessionmaker | Engine,
tenant_id: str,
user: Account | EndUser,
app_id: str | None,
triggered_from: WorkflowRunTriggeredFrom | None,
@@ -36,7 +36,8 @@ class LogstoreWorkflowExecutionRepository(WorkflowExecutionRepository):
Args:
session_factory: SQLAlchemy sessionmaker or engine for creating sessions
user: Account or EndUser object containing tenant_id, user ID, and role information
tenant_id: Tenant that owns the workflow execution
user: Account or EndUser used for creator attribution
app_id: App ID for filtering by application (can be None)
triggered_from: Source of the execution trigger (DEBUGGING or APP_RUN)
"""
@@ -47,10 +48,8 @@ class LogstoreWorkflowExecutionRepository(WorkflowExecutionRepository):
# Note: Project/logstore/index initialization is done at app startup via ext_logstore
self.logstore_client = AliyunLogStore()
# Extract tenant_id from user
tenant_id = extract_tenant_id(user)
if not tenant_id:
raise ValueError("User must have a tenant_id or current_tenant_id")
raise ValueError("tenant_id is required")
self._tenant_id = tenant_id
# Store app context
@@ -64,7 +63,13 @@ class LogstoreWorkflowExecutionRepository(WorkflowExecutionRepository):
self._creator_user_role = CreatorUserRole.ACCOUNT if isinstance(user, Account) else CreatorUserRole.END_USER
# Initialize SQL repository for dual-write support
self.sql_repository = SQLAlchemyWorkflowExecutionRepository(session_factory, user, app_id, triggered_from)
self.sql_repository = SQLAlchemyWorkflowExecutionRepository(
session_factory=session_factory,
tenant_id=tenant_id,
user=user,
app_id=app_id,
triggered_from=triggered_from,
)
# Control flag for dual-write (write to both LogStore and SQL database)
# Set to True to enable dual-write for safe migration, False to use LogStore only
@@ -26,7 +26,6 @@ from graphon.entities import WorkflowNodeExecution
from graphon.enums import WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus
from graphon.model_runtime.utils.encoders import jsonable_encoder
from graphon.workflow_type_encoder import WorkflowRuntimeTypeConverter
from libs.helper import extract_tenant_id
from models import (
Account,
CreatorUserRole,
@@ -109,6 +108,7 @@ class LogstoreWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository):
def __init__(
self,
session_factory: sessionmaker | Engine,
tenant_id: str,
user: Account | EndUser,
app_id: str | None,
triggered_from: WorkflowNodeExecutionTriggeredFrom | None,
@@ -118,7 +118,8 @@ class LogstoreWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository):
Args:
session_factory: SQLAlchemy sessionmaker or engine for creating sessions
user: Account or EndUser object containing tenant_id, user ID, and role information
tenant_id: Tenant that owns the workflow node execution
user: Account or EndUser used for creator attribution
app_id: App ID for filtering by application (can be None)
triggered_from: Source of the execution trigger (SINGLE_STEP or WORKFLOW_RUN)
"""
@@ -128,10 +129,8 @@ class LogstoreWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository):
# Initialize LogStore client
self.logstore_client = AliyunLogStore()
# Extract tenant_id from user
tenant_id = extract_tenant_id(user)
if not tenant_id:
raise ValueError("User must have a tenant_id or current_tenant_id")
raise ValueError("tenant_id is required")
self._tenant_id = tenant_id
# Store app context
@@ -145,7 +144,13 @@ class LogstoreWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository):
self._creator_user_role = CreatorUserRole.ACCOUNT if isinstance(user, Account) else CreatorUserRole.END_USER
# Initialize SQL repository for dual-write support
self.sql_repository = SQLAlchemyWorkflowNodeExecutionRepository(session_factory, user, app_id, triggered_from)
self.sql_repository = SQLAlchemyWorkflowNodeExecutionRepository(
session_factory=session_factory,
tenant_id=tenant_id,
user=user,
app_id=app_id,
triggered_from=triggered_from,
)
# Control flag for dual-write (write to both LogStore and SQL database)
# Set to True to enable dual-write for safe migration, False to use LogStore only
+8 -3
View File
@@ -69,12 +69,17 @@ def on_user_loaded(_sender, user: Union["Account", "EndUser"]):
if user:
try:
current_span = get_current_span()
if not current_span.is_recording():
return
tenant_id = extract_tenant_id(user)
if not tenant_id:
return
if current_span:
current_span.set_attribute(DifySpanAttributes.TENANT_ID, tenant_id)
current_span.set_attribute(GenAIAttributes.USER_ID, user.id)
current_span.set_attributes(
{
DifySpanAttributes.TENANT_ID: tenant_id,
GenAIAttributes.USER_ID: user.id,
}
)
except Exception:
logger.exception("Error setting tenant and user attributes")
pass
+16 -1
View File
@@ -67,5 +67,20 @@ def preserve_flask_contexts(
pass
def set_login_user(user: "Account | EndUser"):
def set_login_user(user: "Account | EndUser") -> None:
"""Set the explicit Flask user and its primitive logging identity snapshot."""
g._login_user = user
from core.logging.context import set_identity_context
from models import Account, EndUser
set_identity_context()
try:
match user:
case Account():
set_identity_context(tenant_id=user.current_tenant_id, user_id=user.id, user_type="account")
case EndUser():
set_identity_context(tenant_id=user.tenant_id, user_id=user.id, user_type=user.type or "end_user")
except Exception:
# Restoring an execution context must not fail because logging metadata is unavailable.
return
+19 -2
View File
@@ -34,6 +34,7 @@ class OAuthState(TypedDict, total=False):
invite_token: str
timezone: str
language: str
redirect_url: str
class GitHubEmailRecord(TypedDict, total=False):
@@ -72,6 +73,7 @@ def encode_oauth_state(
invite_token: str | None = None,
timezone: str | None = None,
language: str | None = None,
redirect_url: str | None = None,
) -> str | None:
state: OAuthState = {}
if invite_token:
@@ -80,6 +82,8 @@ def encode_oauth_state(
state["timezone"] = timezone
if language:
state["language"] = language
if redirect_url:
state["redirect_url"] = redirect_url
if not state:
return None
@@ -122,6 +126,7 @@ class OAuth:
invite_token: str | None = None,
timezone: str | None = None,
language: str | None = None,
redirect_url: str | None = None,
) -> str:
raise NotImplementedError()
@@ -151,13 +156,19 @@ class GitHubOAuth(OAuth):
invite_token: str | None = None,
timezone: str | None = None,
language: str | None = None,
redirect_url: str | None = None,
) -> str:
params = {
"client_id": self.client_id,
"redirect_uri": self.redirect_uri,
"scope": "user:email", # Request only basic user information
}
state = encode_oauth_state(invite_token=invite_token, timezone=timezone, language=language)
state = encode_oauth_state(
invite_token=invite_token,
timezone=timezone,
language=language,
redirect_url=redirect_url,
)
if state:
params["state"] = state
return f"{self._AUTH_URL}?{urllib.parse.urlencode(params)}"
@@ -248,6 +259,7 @@ class GoogleOAuth(OAuth):
invite_token: str | None = None,
timezone: str | None = None,
language: str | None = None,
redirect_url: str | None = None,
) -> str:
params = {
"client_id": self.client_id,
@@ -255,7 +267,12 @@ class GoogleOAuth(OAuth):
"redirect_uri": self.redirect_uri,
"scope": "openid email",
}
state = encode_oauth_state(invite_token=invite_token, timezone=timezone, language=language)
state = encode_oauth_state(
invite_token=invite_token,
timezone=timezone,
language=language,
redirect_url=redirect_url,
)
if state:
params["state"] = state
return f"{self._AUTH_URL}?{urllib.parse.urlencode(params)}"
+12
View File
@@ -214,6 +214,18 @@ class EndUserType(StrEnum):
SERVICE_API = "service-api"
TRIGGER = "trigger"
@classmethod
@override
def _missing_(cls, value):
# Legacy rows persisted the service-api type with an underscore before it
# was normalized to the hyphenated value. The
# `4f7b2c8d9a10_normalize_legacy_end_user_type` migration rewrites those
# rows, but tolerate the old value here as well so an unmigrated end user
# keeps loading instead of failing enum validation on every request.
if value == "service_api":
return cls.SERVICE_API
return super()._missing_(value)
class DocumentDocType(StrEnum):
"""Document doc_type classification"""

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