Compare commits

..
Author SHA1 Message Date
crazywoola fe676f6797 fix missing answer workflow preview 2026-07-21 10:49:09 +08:00
1008 changed files with 15282 additions and 37659 deletions
@@ -32,11 +32,12 @@ Keep this skill focused on Cucumber, Playwright, and package-level E2E guidance.
- `e2e/` uses Cucumber for scenarios and Playwright as the browser layer.
- `DifyWorld` is the per-scenario context object. Type `this` as `DifyWorld` and use `async function`, not arrow functions.
- Keep glue organized by capability under `e2e/features/step-definitions/`; use `common/` only for broadly reusable steps.
- Treat `e2e/AGENTS.md`, `features/support/hooks.ts`, and the Cucumber configuration as the owners of current session and tag semantics. Verify them when behavior depends on session state instead of copying a tag inventory into this skill.
- Browser session behavior comes from `features/support/hooks.ts`:
- default: authenticated session with shared storage state
- `@unauthenticated`: clean browser context
- `@authenticated`: readability/selective-run tag only unless implementation changes
- `@fresh`: only for `e2e:full*` flows
- Do not import Playwright Test runner patterns that bypass the current Cucumber + `DifyWorld` architecture unless the task is explicitly about changing that architecture.
- Perform the behavior under test through Playwright. APIs are allowed for setup, seed preparation, persistence polling, and cleanup, but ordinary Console JSON and representable multipart operations must use the scenario- or process-owned generated oRPC client with request and response validation enabled. Keep the setup/cleanup API identity independent from an unauthenticated or logged-out behavior browser.
- Consume generated operations directly. Do not add one-to-one API wrappers, handwritten endpoint URLs, response DTO casts, duplicate schemas, global mutable clients, or TanStack Query caching in Cucumber. Keep helpers only for real fixture construction, multi-operation orchestration, invariants, polling, derived test views, or protocol adapters.
- Keep SSE, binary, redirect-only, external-service, and readiness exceptions centralized under their protocol owner. A contract mismatch must fail and be fixed at the backend schema owner followed by regeneration; never weaken validation to make E2E pass.
## Workflow
@@ -65,7 +66,7 @@ Keep this skill focused on Cucumber, Playwright, and package-level E2E guidance.
- If a product element has real user-facing semantics but no accessible name, prefer fixing that accessible contract over adding a test id.
5. Validate narrowly.
- Run the narrowest tagged scenario or flow that exercises the change.
- Run the package-required static checks documented in `e2e/AGENTS.md`.
- Run `vpr lint --fix --quiet` from the repository root and `pnpm -C e2e type-check`.
- Broaden verification only when the change affects hooks, tags, setup, or shared step semantics.
## Review Checklist
@@ -76,8 +77,6 @@ Keep this skill focused on Cucumber, Playwright, and package-level E2E guidance.
- Are locators user-facing and assertions web-first?
- Does the change introduce hidden coupling across scenarios, tags, or instance state?
- Does it document or implement behavior that differs from the real hooks or configuration?
- Does setup/cleanup use the generated client directly, with any remaining helper owning more than a one-to-one endpoint forward?
- Is every raw HTTP call a documented protocol or infrastructure exception rather than an ordinary Console operation?
Lead findings with correctness, flake risk, and architecture drift.
+44 -17
View File
@@ -8,6 +8,7 @@
# Lint bulk suppression baselines.
/oxlint-suppressions.json
/eslint-suppressions.json
# CODEOWNERS file
/.github/CODEOWNERS @laipz8200 @crazywoola
@@ -32,9 +33,31 @@
# Backend (default owner, more specific rules below will override)
/api/ @QuantumGhost
# Backend - MCP
/api/core/mcp/ @Nov1c444
/api/core/entities/mcp_provider.py @Nov1c444
/api/services/tools/mcp_tools_manage_service.py @Nov1c444
/api/controllers/mcp/ @Nov1c444
/api/controllers/console/app/mcp_server.py @Nov1c444
# Backend - Tests
/api/tests/ @laipz8200 @QuantumGhost
/api/tests/**/*mcp* @Nov1c444
# Backend - Workflow - Engine (Core graph execution engine)
/api/core/workflow/graph_engine/ @laipz8200 @QuantumGhost
/api/core/workflow/runtime/ @laipz8200 @QuantumGhost
/api/core/workflow/graph/ @laipz8200 @QuantumGhost
/api/core/workflow/graph_events/ @laipz8200 @QuantumGhost
/api/core/workflow/node_events/ @laipz8200 @QuantumGhost
# Backend - Workflow - Nodes (Agent, Iteration, Loop, LLM)
/api/core/workflow/nodes/agent/ @Nov1c444
/api/core/workflow/nodes/iteration/ @Nov1c444
/api/core/workflow/nodes/loop/ @Nov1c444
/api/core/workflow/nodes/llm/ @Nov1c444
# Backend - RAG (Retrieval Augmented Generation)
/api/core/rag/ @JohnJyong
/api/services/rag_pipeline/ @JohnJyong
@@ -88,6 +111,7 @@
/api/core/app/layers/trigger_post_layer.py @CourTeous33
/api/services/trigger/ @CourTeous33
/api/models/trigger.py @CourTeous33
/api/fields/workflow_trigger_fields.py @CourTeous33
/api/repositories/workflow_trigger_log_repository.py @CourTeous33
/api/repositories/sqlalchemy_workflow_trigger_log_repository.py @CourTeous33
/api/libs/schedule_utils.py @CourTeous33
@@ -112,11 +136,11 @@
/api/controllers/console/billing/ @hj24 @zyssyz123
# Backend - Enterprise
/api/configs/enterprise/ @GareArc
/api/services/enterprise/ @GareArc
/api/services/feature_service.py @GareArc
/api/controllers/console/feature.py @GareArc
/api/controllers/web/feature.py @GareArc
/api/configs/enterprise/ @GarfieldDai @GareArc
/api/services/enterprise/ @GarfieldDai @GareArc
/api/services/feature_service.py @GarfieldDai @GareArc
/api/controllers/console/feature.py @GarfieldDai @GareArc
/api/controllers/web/feature.py @GarfieldDai @GareArc
# Backend - Database Migrations
/api/migrations/ @snakevash @laipz8200 @MRZHUH
@@ -129,6 +153,7 @@
# Frontend - Platform and Features
/web/config/ @lyzno1
/web/contract/ @lyzno1
/web/env.ts @lyzno1
/web/features/ @lyzno1
/web/hooks/ @lyzno1
@@ -187,6 +212,7 @@
/web/app/components/rag-pipeline/store/ @iamjoel @zxhlyh
# Frontend - RAG - Documents List
/web/app/components/datasets/documents/list.tsx @iamjoel @WTW0313
/web/app/components/datasets/documents/create-from-pipeline/ @iamjoel @WTW0313
# Frontend - RAG - Segments List
@@ -205,22 +231,22 @@
/web/app/components/plugins/marketplace/ @iamjoel @Yessenia-d
# Frontend - Login and Registration
/web/app/signin/ @iamjoel
/web/app/signup/ @iamjoel
/web/app/reset-password/ @iamjoel
/web/app/install/ @iamjoel
/web/app/init/ @iamjoel
/web/app/forgot-password/ @iamjoel
/web/app/account/ @iamjoel
/web/app/signin/ @douxc @iamjoel
/web/app/signup/ @douxc @iamjoel
/web/app/reset-password/ @douxc @iamjoel
/web/app/install/ @douxc @iamjoel
/web/app/init/ @douxc @iamjoel
/web/app/forgot-password/ @douxc @iamjoel
/web/app/account/ @douxc @iamjoel
# Frontend - Service Authentication
/web/service/base.ts @iamjoel
/web/service/base.ts @douxc @iamjoel
# Frontend - WebApp Authentication and Access Control
/web/app/(shareLayout)/components/ @iamjoel
/web/app/(shareLayout)/webapp-signin/ @iamjoel
/web/app/(shareLayout)/webapp-reset-password/ @iamjoel
/web/app/components/app/app-access-control/ @iamjoel
/web/app/(shareLayout)/components/ @douxc @iamjoel
/web/app/(shareLayout)/webapp-signin/ @douxc @iamjoel
/web/app/(shareLayout)/webapp-reset-password/ @douxc @iamjoel
/web/app/components/app/app-access-control/ @douxc @iamjoel
# Frontend - Explore Page
/web/app/components/explore/ @CodingOnStar @iamjoel
@@ -239,6 +265,7 @@
/web/app/components/base/**/*.spec.tsx @hyoban @CodingOnStar
# Frontend - Utils and Hooks
/web/utils/classnames.ts @iamjoel @zxhlyh
/web/utils/time.ts @iamjoel @zxhlyh
/web/utils/format.ts @iamjoel @zxhlyh
/web/utils/clipboard.ts @iamjoel @zxhlyh
-2
View File
@@ -335,8 +335,6 @@ jobs:
- check-changes
if: needs.pre_job.outputs.should_skip != 'true' && needs.check-changes.outputs.e2e-changed == 'true'
uses: ./.github/workflows/web-e2e.yml
with:
run-external-runtime: false
secrets: inherit
web-e2e-skip:
-16
View File
@@ -26,30 +26,19 @@ jobs:
external_e2e:
- 'e2e/features/agent-v2/**'
- 'e2e/features/step-definitions/agent-v2/**'
- 'e2e/features/step-definitions/common/**'
- 'e2e/features/support/**'
- 'e2e/fixtures/auth.ts'
- 'e2e/fixtures/test-materials/**'
- 'e2e/scripts/**'
- 'e2e/support/**'
- 'e2e/cucumber.config.ts'
- 'e2e/package.json'
- 'e2e/test-env.ts'
- 'e2e/tsconfig.json'
- 'e2e/tsx-register.js'
- 'package.json'
- 'pnpm-lock.yaml'
- '.nvmrc'
- '.github/workflows/post-merge.yml'
- '.github/workflows/web-e2e.yml'
- '.github/actions/setup-web/**'
- 'docker/docker-compose.middleware.yaml'
- 'docker/envs/middleware.env.example'
- 'dify-agent/**'
- 'dify-agent-runtime/**'
- 'api/pyproject.toml'
- 'api/uv.lock'
- 'api/tests/integration_tests/.env.example'
- 'api/clients/agent_backend/**'
- 'api/core/app/apps/agent_app/**'
- 'api/core/workflow/nodes/agent_v2/**'
@@ -59,13 +48,8 @@ jobs:
- 'api/services/plugin/**'
- 'api/core/tools/**'
- 'api/services/tools/**'
- 'packages/contracts/package.json'
- 'packages/contracts/generated/api/console/agent/**'
- 'packages/contracts/generated/api/console/apps/**'
- 'packages/contracts/generated/api/console/datasets/**'
- 'packages/contracts/generated/api/console/orpc.gen.ts'
- 'packages/contracts/generated/api/console/workspaces/**'
- 'packages/contracts/generated/api/service/**'
- 'web/features/agent-v2/**'
- 'web/app/(commonLayout)/agents/**'
- 'web/app/(commonLayout)/@detailSidebar/agents/**'
+4 -7
View File
@@ -4,9 +4,9 @@ on:
workflow_call:
inputs:
run-external-runtime:
description: Run only the prepared and external runtime suite instead of the core suites.
required: true
required: false
type: boolean
default: false
permissions:
contents: read
@@ -46,7 +46,6 @@ jobs:
run: uv sync --project api --dev
- name: Run E2E support unit tests
if: ${{ !inputs.run-external-runtime }}
working-directory: ./e2e
run: vp run test:unit
@@ -55,7 +54,6 @@ jobs:
run: vp run e2e:install
- name: Run isolated source-api and built-web Cucumber E2E tests
if: ${{ !inputs.run-external-runtime }}
working-directory: ./e2e
env:
E2E_ADMIN_EMAIL: e2e-admin@example.com
@@ -66,7 +64,7 @@ jobs:
run: vp run e2e:full
- name: Preserve Chromium E2E report and logs
if: ${{ !cancelled() && !inputs.run-external-runtime }}
if: ${{ !cancelled() }}
run: |
if [[ -d e2e/cucumber-report ]]; then
mv e2e/cucumber-report e2e/cucumber-report-non-external
@@ -76,7 +74,6 @@ jobs:
fi
- name: Run WebKit keyboard and browser smoke tests
if: ${{ !inputs.run-external-runtime }}
working-directory: ./e2e
env:
E2E_ADMIN_EMAIL: e2e-admin@example.com
@@ -102,7 +99,7 @@ jobs:
vp run e2e -- --tags '@browser-smoke'
- name: Preserve WebKit E2E report and logs
if: ${{ !cancelled() && !inputs.run-external-runtime }}
if: ${{ !cancelled() }}
run: |
if [[ -d e2e/cucumber-report ]]; then
mv e2e/cucumber-report e2e/cucumber-report-webkit
+1 -1
View File
@@ -71,7 +71,7 @@ Dify is an open-source LLM app development platform. Its intuitive interface com
<br/>
The easiest way to start the Dify server is through [Docker Compose](docker/docker-compose.yaml). Before running Dify with the following commands, make sure that [Docker](https://docs.docker.com/get-docker/) and Docker Compose v2.24.0 or later are installed on your machine:
The easiest way to start the Dify server is through [Docker Compose](docker/docker-compose.yaml). Before running Dify with the following commands, make sure that [Docker](https://docs.docker.com/get-docker/) and [Docker Compose](https://docs.docker.com/compose/install/) are installed on your machine:
```bash
cd dify
-6
View File
@@ -14,12 +14,6 @@ class EnterpriseFeatureConfig(BaseSettings):
default=False,
)
WEBAPP_PUBLIC_ACCESS_ENABLED: bool = Field(
description="Whether admins are allowed to set a webapp's access mode to public (anyone with the link, "
"no auth). Disable in security-sensitive on-prem deployments.",
default=True,
)
CAN_REPLACE_LOGO: bool = Field(
description="Allow customization of the enterprise logo.",
default=False,
-5
View File
@@ -1497,11 +1497,6 @@ class LoginConfig(BaseSettings):
class AccountConfig(BaseSettings):
ENABLE_CHANGE_EMAIL: bool = Field(
description="whether users can change their email address",
default=True,
)
ACCOUNT_DELETION_TOKEN_EXPIRY_MINUTES: PositiveInt = Field(
description="Duration in minutes for which a account deletion token remains valid",
default=5,
+7 -20
View File
@@ -10,7 +10,6 @@ from extensions.ext_database import db
from libs.login import current_account_with_tenant
from models.dataset import Dataset
from models.model import App
from services.agent.roster_service import AgentRosterService
from services.enterprise.rbac_service import RBACService
__all__ = ["RBACPermission", "RBACResourceScope", "enforce_rbac_access", "rbac_permission_required"]
@@ -52,7 +51,7 @@ def enforce_rbac_access(
check_resource_type = None if resource_type == RBACResourceScope.WORKSPACE else resource_type
resource_id = None
if resource_required and check_resource_type:
resource_id = _extract_resource_id(resource_type, tenant_id, path_args)
resource_id = _extract_resource_id(resource_type, path_args)
if _is_resource_owned_by_current_user(tenant_id, account_id, resource_type, resource_id):
return
allowed = RBACService.CheckAccess.check(
@@ -132,14 +131,11 @@ def _is_resource_owned_by_current_user(
return False
def _extract_resource_id(
resource_type: RBACResourceScope, tenant_id: str, path_args: dict[str, object] | None = None
) -> str:
def _extract_resource_id(resource_type: RBACResourceScope, path_args: dict[str, object] | None = None) -> str:
"""Extract the resource ID from matched path arguments.
Some legacy route classes use neutral names such as ``resource_id`` for
app/dataset resources, and Agent routes carry ``agent_id``, which is
resolved to the App backing that Agent.
app/dataset resources, and Agent App routes use ``agent_id`` as the app id.
Dataset endpoints behind a rag-pipeline route contain ``pipeline_id``
instead of ``dataset_id``. In that case we look up the associated
``Dataset`` row via ``Dataset.pipeline_id``.
@@ -150,19 +146,10 @@ def _extract_resource_id(
matched_args = {**view_args, **(path_args or {})}
if resource_type == RBACResourceScope.APP:
app_id = matched_args.get("app_id")
if app_id:
return str(app_id)
agent_id = matched_args.get("agent_id")
if agent_id:
authz_app_id = AgentRosterService(db.session).peek_authz_app_id(tenant_id=tenant_id, agent_id=str(agent_id))
return authz_app_id or str(agent_id)
resource_id = matched_args.get("resource_id")
if resource_id:
return str(resource_id)
raise ValueError("Missing app_id in request path")
app_id = matched_args.get("app_id") or matched_args.get("agent_id") or matched_args.get("resource_id")
if not app_id:
raise ValueError("Missing app_id in request path")
return str(app_id)
if resource_type == RBACResourceScope.DATASET:
dataset_id = matched_args.get("dataset_id") or matched_args.get("resource_id")
@@ -230,7 +230,6 @@ class WorkflowAgentComposerSaveToRosterApi(Resource):
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT)
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user_id
@with_current_tenant_id
@with_session
@@ -440,7 +439,6 @@ class SnippetAgentComposerSaveToRosterApi(Resource):
@rbac_permission_required(
RBACResourceScope.WORKSPACE, RBACPermission.SNIPPETS_CREATE_AND_MODIFY, resource_required=False
)
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user_id
@with_current_tenant_id
@with_session
@@ -480,7 +478,6 @@ class AgentComposerApi(Resource):
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT)
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user_id
@with_current_tenant_id
@with_session
+1 -44
View File
@@ -62,7 +62,7 @@ from libs.datetime_utils import parse_time_range
from libs.helper import dump_response
from libs.login import login_required
from models import Account
from models.agent import Agent, AgentConfigDraftType, AgentStatus
from models.agent import Agent, AgentStatus
from models.agent_config_entities import AgentSoulConfig
from models.enums import ApiTokenType
from models.model import ApiToken, App, IconType
@@ -266,13 +266,6 @@ class AgentDebugConversationRefreshResponse(BaseModel):
debug_conversation_message_count: int = 0
class AgentDebugConversationRefreshPayload(BaseModel):
draft_type: AgentConfigDraftType = Field(
default=AgentConfigDraftType.DEBUG_BUILD,
description="Agent draft surface whose conversation should be refreshed",
)
class AgentPublishPayload(BaseModel):
version_note: str | None = Field(default=None, description="Optional note for this published Agent version")
@@ -316,7 +309,6 @@ register_schema_models(
AgentAppCopyPayload,
AgentPublishPayload,
AgentBuildDraftCheckoutPayload,
AgentDebugConversationRefreshPayload,
ComposerSavePayload,
AgentApiStatusPayload,
AgentInviteOptionsQuery,
@@ -400,7 +392,6 @@ def _serialize_agent_app_detail(
tenant_id=app_model.tenant_id,
agent_id=agent.id,
account_id=current_user.id,
draft_type=AgentConfigDraftType.DEBUG_BUILD,
commit=False,
)
message_count = roster_service.count_agent_app_debug_conversation_messages(
@@ -448,7 +439,6 @@ def _serialize_agent_app_pagination(session: Session, app_pagination, *, tenant_
tenant_id=tenant_id,
agents=list(agents_by_app_id.values()),
account_id=current_user.id,
draft_type=AgentConfigDraftType.DEBUG_BUILD,
)
payload = AgentAppPagination.model_validate(
app_pagination,
@@ -552,7 +542,6 @@ class AgentAppListApi(Resource):
@setup_required
@login_required
@account_initialization_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session
@@ -590,7 +579,6 @@ class AgentAppListApi(Resource):
@login_required
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session
@@ -632,7 +620,6 @@ class AgentAppApi(Resource):
@login_required
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session
@@ -658,7 +645,6 @@ class AgentAppApi(Resource):
@login_required
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_tenant_id
@with_session
def delete(self, session: Session, tenant_id: str, agent_id: UUID):
@@ -669,16 +655,6 @@ class AgentAppApi(Resource):
@console_ns.route("/agent/<uuid:agent_id>/debug-conversation/refresh")
class AgentDebugConversationRefreshApi(Resource):
@console_ns.expect(console_ns.models[AgentDebugConversationRefreshPayload.__name__])
@console_ns.doc(
params={
"payload": {
"in": "body",
"required": False,
"schema": {"$ref": f"#/components/schemas/{AgentDebugConversationRefreshPayload.__name__}"},
}
}
)
@console_ns.response(
200,
"Agent debug conversation refreshed",
@@ -693,12 +669,10 @@ class AgentDebugConversationRefreshApi(Resource):
@with_current_tenant_id
@with_session
def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID):
args = AgentDebugConversationRefreshPayload.model_validate(request.get_json(silent=True) or {})
debug_conversation_id = _agent_roster_service(session).refresh_agent_app_debug_conversation_id(
tenant_id=tenant_id,
agent_id=str(agent_id),
account_id=current_user.id,
draft_type=args.draft_type,
)
return AgentDebugConversationRefreshResponse(
debug_conversation_id=debug_conversation_id,
@@ -716,7 +690,6 @@ class AgentPublishApi(Resource):
@login_required
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session
@@ -739,7 +712,6 @@ class AgentBuildDraftCheckoutApi(Resource):
@login_required
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session
@@ -757,7 +729,6 @@ class AgentBuildDraftCheckoutApi(Resource):
@console_ns.route("/agent/<uuid:agent_id>/build-draft")
class AgentBuildDraftApi(Resource):
@console_ns.response(200, "Agent build draft", console_ns.models[AgentBuildDraftResponse.__name__])
@console_ns.response(404, "Agent build draft not found")
@setup_required
@login_required
@account_initialization_required
@@ -816,7 +787,6 @@ class AgentBuildDraftApplyApi(Resource):
@login_required
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session
@@ -839,7 +809,6 @@ class AgentAppCopyApi(Resource):
@login_required
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session
@@ -865,7 +834,6 @@ class AgentApiAccessApi(Resource):
@setup_required
@login_required
@account_initialization_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_tenant_id
@with_session(write=False)
def get(self, session: Session, tenant_id: str, agent_id: UUID):
@@ -882,7 +850,6 @@ class AgentApiStatusApi(Resource):
@login_required
@is_admin_or_owner_required
@account_initialization_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION)
@with_current_tenant_id
@with_session
@@ -901,7 +868,6 @@ class AgentApiKeyListApi(BaseApiKeyListResource):
token_prefix = "app-"
@console_ns.response(200, "Agent service API keys", console_ns.models[ApiKeyList.__name__])
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_tenant_id
@with_session(write=False)
def get(self, session: Session, tenant_id: str, agent_id: UUID) -> dict[str, object]:
@@ -912,7 +878,6 @@ class AgentApiKeyListApi(BaseApiKeyListResource):
@console_ns.response(400, "Maximum keys exceeded")
@with_current_tenant_id
@edit_permission_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION)
@with_session
def post(self, session: Session, tenant_id: str, agent_id: UUID) -> tuple[dict[str, object], int]:
@@ -932,7 +897,6 @@ class AgentApiKeyApi(BaseApiKeyResource):
@console_ns.response(204, "Agent service API key deleted")
@with_current_user
@with_current_tenant_id
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION)
@with_session
def delete(
@@ -978,7 +942,6 @@ class AgentLogsApi(Resource):
@setup_required
@login_required
@account_initialization_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session(write=False)
@@ -1017,7 +980,6 @@ class AgentLogMessagesApi(Resource):
@setup_required
@login_required
@account_initialization_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session(write=False)
@@ -1056,7 +1018,6 @@ class AgentLogSourcesApi(Resource):
@setup_required
@login_required
@account_initialization_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session(write=False)
@@ -1077,7 +1038,6 @@ class AgentStatisticsSummaryApi(Resource):
@setup_required
@login_required
@account_initialization_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session(write=False)
@@ -1103,7 +1063,6 @@ class AgentRosterVersionsApi(Resource):
@setup_required
@login_required
@account_initialization_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_tenant_id
@with_session(write=False)
def get(self, session: Session, tenant_id: str, agent_id: UUID):
@@ -1119,7 +1078,6 @@ class AgentRosterVersionDetailApi(Resource):
@setup_required
@login_required
@account_initialization_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_tenant_id
@with_session(write=False)
def get(self, session: Session, tenant_id: str, agent_id: UUID, version_id: UUID):
@@ -1140,7 +1098,6 @@ class AgentRosterVersionRestoreApi(Resource):
@login_required
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session
-4
View File
@@ -12,7 +12,6 @@ from werkzeug.exceptions import Forbidden
from configs import dify_config
from controllers.common.schema import register_response_schema_models
from controllers.common.session import with_session
from controllers.console.app.wraps import agent_manage_required_for_agent_app
from fields.base import ResponseModel
from libs.helper import dump_response, to_timestamp
from libs.login import login_required
@@ -195,7 +194,6 @@ class AppApiKeyListResource(BaseApiKeyListResource):
@console_ns.doc(params={"resource_id": "App ID"})
@console_ns.response(200, "API keys retrieved successfully", console_ns.models[ApiKeyList.__name__])
@with_current_tenant_id
@agent_manage_required_for_agent_app
@with_session(write=False)
def get(self, session: Session, current_tenant_id: str, resource_id: UUID) -> dict[str, object]:
"""Get all API keys for an app"""
@@ -212,7 +210,6 @@ class AppApiKeyListResource(BaseApiKeyListResource):
@with_current_tenant_id
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION)
@agent_manage_required_for_agent_app
@with_session
def post(self, session: Session, current_tenant_id: str, resource_id: UUID) -> tuple[dict[str, object], int]:
"""Create a new API key for an app"""
@@ -236,7 +233,6 @@ class AppApiKeyResource(BaseApiKeyResource):
@with_current_user
@with_current_tenant_id
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION)
@agent_manage_required_for_agent_app
@with_session
def delete(
self,
+17 -34
View File
@@ -9,7 +9,7 @@ from flask_restx import Resource
from pydantic import AliasChoices, BaseModel, Field, ValidationInfo, computed_field, field_validator, model_validator
from sqlalchemy import select
from sqlalchemy.orm import Session
from werkzeug.exceptions import BadRequest, Forbidden, NotFound
from werkzeug.exceptions import BadRequest, NotFound
from configs import dify_config
from controllers.common.app_access import resolve_app_access_filter
@@ -23,7 +23,7 @@ from controllers.common.schema import (
register_schema_models,
)
from controllers.console import console_ns
from controllers.console.app.wraps import agent_manage_required_for_agent_app, get_app_model, with_session
from controllers.console.app.wraps import get_app_model, with_session
from controllers.console.workspace.models import LoadBalancingPayload
from controllers.console.wraps import (
RBACPermission,
@@ -75,7 +75,6 @@ from services.entities.knowledge_entities.knowledge_entities import (
WeightModel,
WeightVectorSetting,
)
from services.errors.account import NoPermissionError
from services.feature_service import FeatureService
from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task
@@ -247,7 +246,7 @@ class ModelConfigPartial(ResponseModel):
return to_timestamp(value)
class AppModelConfigResponse(ResponseModel):
class ModelConfig(ResponseModel):
opening_statement: str | None = None
suggested_questions: Any | None = Field(
default=None, validation_alias=AliasChoices("suggested_questions_list", "suggested_questions")
@@ -420,7 +419,7 @@ class AppDetail(AppResponseModel):
icon_background: str | None = None
enable_site: bool
enable_api: bool
model_config_: AppModelConfigResponse | None = Field(
model_config_: ModelConfig | None = Field(
default=None,
validation_alias=AliasChoices("app_model_config", "model_config"),
alias="model_config",
@@ -526,13 +525,7 @@ def _enrich_app_list_items(session: Session, *, apps: Sequence[App], tenant_id:
register_enum_models(console_ns, RetrievalMethod, WorkflowExecutionStatus, DatasetPermissionEnum)
register_response_schema_models(
console_ns,
RedirectUrlResponse,
SimpleResultResponse,
AppImportResponse,
AppTraceResponse,
AppModelConfigResponse,
AppDetail,
console_ns, RedirectUrlResponse, SimpleResultResponse, AppImportResponse, AppTraceResponse
)
register_schema_models(
@@ -551,8 +544,10 @@ register_schema_models(
Tag,
WorkflowPartial,
ModelConfigPartial,
ModelConfig,
AppDetailSiteResponse,
DeletedTool,
AppDetail,
AppExportResponse,
Segmentation,
PreProcessingRule,
@@ -828,7 +823,6 @@ class AppApi(Resource):
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT)
@agent_manage_required_for_agent_app
@with_session
@get_app_model(mode=None)
def put(self, session: Session, app_model: App):
@@ -863,7 +857,6 @@ class AppApi(Resource):
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_DELETE)
@agent_manage_required_for_agent_app
@with_session
@get_app_model
def delete(self, session: Session, app_model: App):
@@ -888,7 +881,6 @@ class AppCopyApi(Resource):
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_CREATE_AND_MANAGEMENT)
@agent_manage_required_for_agent_app
@with_current_user
@with_current_tenant_id
@get_app_model(mode=None)
@@ -900,19 +892,16 @@ class AppCopyApi(Resource):
with Session(db.engine, expire_on_commit=False) as session:
import_service = AppDslService(session)
yaml_content = import_service.export_dsl(app_model=app_model, session=session, include_secret=True)
try:
result = import_service.import_app(
account=current_user,
import_mode=ImportMode.YAML_CONTENT,
yaml_content=yaml_content,
name=args.name,
description=args.description,
icon_type=args.icon_type,
icon=args.icon,
icon_background=args.icon_background,
)
except NoPermissionError as e:
raise Forbidden(str(e))
result = import_service.import_app(
account=current_user,
import_mode=ImportMode.YAML_CONTENT,
yaml_content=yaml_content,
name=args.name,
description=args.description,
icon_type=args.icon_type,
icon=args.icon,
icon_background=args.icon_background,
)
if result.status == ImportStatus.FAILED:
session.rollback()
return dump_response(AppImportResponse, result), 400
@@ -966,7 +955,6 @@ class AppExportApi(Resource):
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_IMPORT_EXPORT_DSL)
@agent_manage_required_for_agent_app
@get_app_model
def get(self, app_model: App):
"""Export app"""
@@ -991,7 +979,6 @@ class AppPublishToCreatorsPlatformApi(Resource):
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_IMPORT_EXPORT_DSL)
@agent_manage_required_for_agent_app
@with_current_user_id
@get_app_model(mode=None)
def post(self, current_user_id: str, app_model: App):
@@ -1022,7 +1009,6 @@ class AppNameApi(Resource):
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT)
@agent_manage_required_for_agent_app
@with_session
@get_app_model(mode=None)
def post(self, session: Session, app_model: App):
@@ -1050,7 +1036,6 @@ class AppIconApi(Resource):
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT)
@agent_manage_required_for_agent_app
@with_session
@get_app_model(mode=None)
def post(self, session: Session, app_model: App):
@@ -1084,7 +1069,6 @@ class AppSiteStatus(Resource):
@account_initialization_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION)
@agent_manage_required_for_agent_app
@with_session
@get_app_model(mode=None)
def post(self, session: Session, app_model: App):
@@ -1112,7 +1096,6 @@ class AppApiStatus(Resource):
@is_admin_or_owner_required
@account_initialization_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION)
@agent_manage_required_for_agent_app
@with_session
@get_app_model(mode=None)
def post(self, session: Session, app_model: App):
+13 -21
View File
@@ -1,7 +1,6 @@
from flask_restx import Resource
from pydantic import BaseModel, Field
from sqlalchemy.orm import Session
from werkzeug.exceptions import Forbidden
from configs import dify_config
from controllers.common.schema import register_enum_models, register_schema_models
@@ -29,7 +28,6 @@ from services.app_dsl_service import (
)
from services.enterprise.enterprise_service import EnterpriseService
from services.entities.dsl_entities import CheckDependenciesResult, ImportStatus
from services.errors.account import NoPermissionError
from services.feature_service import FeatureService
from .. import console_ns
@@ -93,21 +91,18 @@ class AppImportApi(Resource):
import_service = AppDslService(session)
# Import app
account = current_user
try:
result = import_service.import_app(
account=account,
import_mode=args.mode,
yaml_content=args.yaml_content,
yaml_url=args.yaml_url,
name=args.name,
description=args.description,
icon_type=args.icon_type,
icon=args.icon,
icon_background=args.icon_background,
app_id=args.app_id,
)
except NoPermissionError as e:
raise Forbidden(str(e))
result = import_service.import_app(
account=account,
import_mode=args.mode,
yaml_content=args.yaml_content,
yaml_url=args.yaml_url,
name=args.name,
description=args.description,
icon_type=args.icon_type,
icon=args.icon,
icon_background=args.icon_background,
app_id=args.app_id,
)
if result.status == ImportStatus.FAILED:
session.rollback()
else:
@@ -162,10 +157,7 @@ class AppImportConfirmApi(Resource):
import_service = AppDslService(session)
# Confirm import
account = current_user
try:
result = import_service.confirm_import(import_id=import_id, account=account)
except NoPermissionError as e:
raise Forbidden(str(e))
result = import_service.confirm_import(import_id=import_id, account=account)
if result.status == ImportStatus.FAILED:
session.rollback()
else:
+1 -14
View File
@@ -49,7 +49,6 @@ from libs import helper
from libs.helper import uuid_value
from libs.login import login_required
from models import Account
from models.agent import AgentConfigDraftType
from models.model import App, AppMode
from services.agent.errors import AgentNotFoundError
from services.agent.roster_service import AgentRosterService
@@ -344,23 +343,14 @@ class AgentChatMessageStopApi(Resource):
def _resolve_current_user_agent_debug_conversation_id(
*,
session: Session,
current_tenant_id: str,
current_user: Account,
app_model: App,
agent_id: str | None,
draft_type: AgentConfigDraftType,
*, session: Session, current_tenant_id: str, current_user: Account, app_model: App, agent_id: str | None
) -> str:
"""Resolve the current editor's conversation without crossing draft surfaces."""
roster_service = AgentRosterService(session)
if agent_id:
return roster_service.get_or_create_agent_app_debug_conversation_id(
tenant_id=current_tenant_id,
agent_id=agent_id,
account_id=current_user.id,
draft_type=draft_type,
)
agent = roster_service.get_app_backing_agent(tenant_id=current_tenant_id, app_id=str(app_model.id))
@@ -370,7 +360,6 @@ def _resolve_current_user_agent_debug_conversation_id(
tenant_id=current_tenant_id,
agent_id=agent.id,
account_id=current_user.id,
draft_type=draft_type,
)
@@ -393,7 +382,6 @@ def _create_chat_message(
current_user=current_user,
app_model=app_model,
agent_id=agent_id,
draft_type=AgentConfigDraftType(args_model.draft_type),
)
if args_model.conversation_id and args_model.conversation_id != debug_conversation_id:
raise NotFound("Conversation Not Exists.")
@@ -430,7 +418,6 @@ def _create_build_chat_finalization_message(
current_user=current_user,
app_model=app_model,
agent_id=agent_id,
draft_type=AgentConfigDraftType.DEBUG_BUILD,
)
args: dict[str, Any] = {
"query": _BUILD_CHAT_FINALIZATION_QUERY,
+1 -3
View File
@@ -10,7 +10,7 @@ from constants.languages import supported_language
from controllers.common.schema import register_schema_models
from controllers.common.session import with_session
from controllers.console import console_ns
from controllers.console.app.wraps import agent_manage_required_for_agent_app, get_app_model
from controllers.console.app.wraps import get_app_model
from controllers.console.wraps import (
RBACPermission,
RBACResourceScope,
@@ -93,7 +93,6 @@ class AppSite(Resource):
@login_required
@edit_permission_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION)
@agent_manage_required_for_agent_app
@account_initialization_required
@with_current_user
@with_session
@@ -146,7 +145,6 @@ class AppSiteAccessTokenReset(Resource):
@login_required
@is_admin_or_owner_required
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION)
@agent_manage_required_for_agent_app
@account_initialization_required
@with_current_user
@with_session
+13 -16
View File
@@ -359,12 +359,6 @@ class WorkflowPublishResponse(ResponseModel):
created_at: int
class SyncDraftWorkflowResponse(ResponseModel):
result: str
hash: str
updated_at: int
class WorkflowRestoreResponse(ResponseModel):
result: str
hash: str
@@ -447,7 +441,6 @@ register_response_schema_models(
WorkflowOnlineUsersByApp,
WorkflowOnlineUsersResponse,
WorkflowPublishResponse,
SyncDraftWorkflowResponse,
WorkflowRestoreResponse,
DefaultBlockConfigsResponse,
DefaultBlockConfigResponse,
@@ -563,7 +556,14 @@ class DraftWorkflowApi(Resource):
@console_ns.response(
200,
"Draft workflow synced successfully",
console_ns.models[SyncDraftWorkflowResponse.__name__],
console_ns.model(
"SyncDraftWorkflowResponse",
{
"result": fields.String,
"hash": fields.String,
"updated_at": fields.String,
},
),
)
@console_ns.response(400, "Invalid workflow configuration")
@console_ns.response(403, "Permission denied")
@@ -618,14 +618,11 @@ class DraftWorkflowApi(Resource):
except VariableError as e:
raise InvalidArgumentError(description=str(e))
return dump_response(
SyncDraftWorkflowResponse,
{
"result": "success",
"hash": workflow.unique_hash,
"updated_at": TimestampField().format(workflow.updated_at or workflow.created_at),
},
)
return {
"result": "success",
"hash": workflow.unique_hash,
"updated_at": TimestampField().format(workflow.updated_at or workflow.created_at),
}
@console_ns.route("/apps/<uuid:app_id>/advanced-chat/workflows/draft/run")
+1 -48
View File
@@ -12,22 +12,14 @@ from typing import cast, overload
from sqlalchemy import select
from sqlalchemy.orm import Session
from configs import dify_config
from controllers.common.session import with_session
from controllers.common.wraps import RBACPermission, RBACResourceScope, enforce_rbac_access
from controllers.console.app.error import AppNotFoundError
from extensions.ext_database import db
from libs.login import current_account_with_tenant
from models import App, AppMode, TrialApp
from models.agent import AgentScope
from services.recommended_app_service import RecommendedAppService
__all__ = [
"agent_manage_required_for_agent_app",
"get_app_model",
"get_app_model_with_trial",
"with_session",
]
__all__ = ["get_app_model", "get_app_model_with_trial", "with_session"]
def _load_app_model(session: Session, app_id: str) -> App | None:
@@ -56,45 +48,6 @@ def _load_app_model_with_trial(session: Session, app_id: str) -> App | None:
return app_model
def agent_manage_required_for_agent_app[**P, R](view: Callable[P, R]) -> Callable[P, R]:
"""Gate generic app management routes that target an Agent App.
A hidden workflow-only backing App only reuses the App runtime and is not
part of the general app management plane, so generic routes reject it
outright. Managing a roster Agent App mutates the roster Agent behind it
(rename/icon sync, archive, API enablement), so it additionally requires
workspace ``agent.manage`` on top of the route's existing App permission
checks when RBAC is enabled. A no-op for non-agent Apps. Must be placed
above ``get_app_model`` so the ``app_id`` path parameter is still present.
"""
@wraps(view)
def decorated(*args: P.args, **kwargs: P.kwargs) -> R:
raw_app_id = kwargs.get("app_id") or kwargs.get("resource_id")
if raw_app_id is not None:
app_model = _load_app_model_from_scoped_session(str(raw_app_id))
binding = (
app_model.agent_app_binding_with_session(session=db.session(), include_archived=True)
if app_model is not None
else None
)
if binding is not None:
if binding.scope == AgentScope.WORKFLOW_ONLY:
raise AppNotFoundError()
if dify_config.RBAC_ENABLED:
current_user, current_tenant_id = current_account_with_tenant()
enforce_rbac_access(
tenant_id=current_tenant_id,
account_id=current_user.id,
resource_type=RBACResourceScope.WORKSPACE,
scene=RBACPermission.AGENT_MANAGE,
resource_required=False,
)
return view(*args, **kwargs)
return decorated
def _get_injected_session(args: tuple[object, ...]) -> Session | None:
"""Return the request session inserted by `with_session`, if this handler has been migrated."""
if len(args) < 2:
-14
View File
@@ -7,13 +7,10 @@ from configs import dify_config
from constants.languages import supported_language
from controllers.common.schema import query_params_from_model, register_schema_models
from controllers.console import console_ns
from controllers.console.auth.error import InvitationAccountMismatchError
from controllers.console.error import AccountInFreezeError, AlreadyActivateError
from extensions.ext_database import db
from libs.datetime_utils import naive_utc_now
from libs.helper import EmailStr, timezone
from libs.login import current_account_with_tenant
from libs.token import extract_access_token
from models import AccountStatus
from models.account import TenantAccountJoin, TenantAccountRole
from services.account_service import RegisterService, TenantService
@@ -139,12 +136,6 @@ class ActivateApi(Resource):
)
@console_ns.response(400, "Already activated or invalid token")
def post(self):
"""Accept an invitation without letting an existing session act for another account.
Token-only activation remains available for legacy clients. When the request already
carries a console session, that session must belong to the account encoded in the
invitation before the token is consumed or tenant membership is changed.
"""
args = ActivatePayload.model_validate(console_ns.payload)
normalized_request_email = args.email.lower() if args.email else None
@@ -155,11 +146,6 @@ class ActivateApi(Resource):
raise AlreadyActivateError()
account = invitation["account"]
if extract_access_token(request):
current_account, _ = current_account_with_tenant()
if current_account.id != account.id:
raise InvitationAccountMismatchError()
if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(account.email):
raise AccountInFreezeError()
-6
View File
@@ -13,12 +13,6 @@ class InvalidEmailError(BaseHTTPException):
code = 400
class InvitationAccountMismatchError(BaseHTTPException):
error_code = "invitation_account_mismatch"
description = "This invitation was sent to another account. Please sign in with the invited account."
code = 403
class PasswordMismatchError(BaseHTTPException):
error_code = "password_mismatch"
description = "The passwords do not match."
+21 -35
View File
@@ -6,7 +6,6 @@ from flask import current_app, redirect, request
from flask_restx import Resource
from pydantic import BaseModel, Field
from werkzeug.exceptions import Unauthorized
from werkzeug.wrappers import Response
from configs import dify_config
from constants.languages import languages
@@ -128,20 +127,6 @@ def _preferred_interface_language(language: str | None = None) -> str:
return languages[0]
def _redirect_with_console_session(account: Account, target_url: str) -> Response:
"""Create a console session and attach its cookies to a redirect response."""
token_pair = AccountService.login(
account=account,
session=db.session(),
ip_address=extract_remote_ip(request),
)
response = redirect(target_url)
set_access_token_to_cookie(request, response, token_pair.access_token)
set_refresh_token_to_cookie(request, response, token_pair.refresh_token)
set_csrf_token_to_cookie(request, response, token_pair.csrf_token)
return response
@console_ns.route("/oauth/login/<provider>")
class OAuthLogin(Resource):
@console_ns.doc("oauth_login")
@@ -210,26 +195,16 @@ class OAuthCallback(Resource):
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message={urllib.parse.quote(str(e))}")
if invite_token and RegisterService.is_valid_invite_token(invite_token):
invitation = RegisterService.get_invitation_if_token_valid(
None,
None,
invite_token,
session=db.session(),
)
if not invitation:
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message=Invalid invitation token.")
if invitation["data"]["email"].lower() != user_info.email.lower():
message = "This invitation was sent to another account. Please sign in with the invited account."
query = urllib.parse.urlencode({"message": message, "invite_token": invite_token})
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?{query}")
invitation = RegisterService.get_invitation_by_token(token=invite_token)
if invitation:
invitation_email = invitation.get("email", None)
invitation_email_normalized = (
invitation_email.lower() if isinstance(invitation_email, str) else invitation_email
)
if invitation_email_normalized != user_info.email.lower():
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message=Invalid invitation token.")
account = invitation["account"]
if account.status == AccountStatus.BANNED:
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message=Account is banned.")
AccountService.link_account_integrate(provider, user_info.id, account, session=db.session())
target_url = f"{dify_config.CONSOLE_WEB_URL}/signin/invite-settings?invite_token={invite_token}"
return _redirect_with_console_session(account, target_url)
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin/invite-settings?invite_token={invite_token}")
try:
account, oauth_new_user = _generate_account(provider, user_info, timezone=timezone, language=language)
@@ -264,10 +239,21 @@ class OAuthCallback(Resource):
"?message=Workspace not found, please contact system admin to invite you to join in a workspace."
)
token_pair = AccountService.login(
account=account,
session=db.session(),
ip_address=extract_remote_ip(request),
)
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()}"
return _redirect_with_console_session(account, target_url)
response = redirect(target_url)
set_access_token_to_cookie(request, response, token_pair.access_token)
set_refresh_token_to_cookie(request, response, token_pair.refresh_token)
set_csrf_token_to_cookie(request, response, token_pair.csrf_token)
return response
def _get_account_by_openid_or_email(provider: str, user_info: OAuthUserInfo) -> Account | None:
+47 -134
View File
@@ -8,9 +8,8 @@ can be validated explicitly against the pinned KnowledgeFS contract during devel
Console auth and contract-specific dataset RBAC run before forwarding. Request
bodies are capped at 64 MiB, JSON and binary responses have separate bounds,
SSE responses remain streaming with a bounded idle read timeout, and only safe
response headers are exposed. Operation-specific upstream error mappings are
applied before Console JSON error handling; the default maps 401 to 502 so it
cannot trigger browser-session recovery and preserves resource-level 403.
response headers are exposed. Upstream 401 responses become 502 so they cannot
trigger Dify browser-session recovery; resource-level 403 responses remain 403.
"""
from __future__ import annotations
@@ -28,11 +27,9 @@ from werkzeug.exceptions import (
BadGateway,
Forbidden,
GatewayTimeout,
HTTPException,
NotFound,
RequestEntityTooLarge,
ServiceUnavailable,
default_exceptions,
)
from configs import dify_config
@@ -44,28 +41,21 @@ from controllers.console.wraps import (
)
from core.helper import ssrf_proxy
from libs.login import current_account_with_tenant, login_required
from services.knowledge_fs_operations import KnowledgeFSMethod
from services.knowledge_fs_proxy import (
KnowledgeFSAccessDeniedError,
KnowledgeFSAuthorization,
KnowledgeFSConfigurationError,
KnowledgeFSMethod,
KnowledgeFSRouteNotAllowedError,
KnowledgeFSTimeoutError,
KnowledgeFSTransportError,
KnowledgeFSUpstreamResponse,
authorize_knowledge_fs_request,
get_knowledge_fs_operation,
proxy_authorized_knowledge_fs_request,
proxy_knowledge_fs_request,
)
logger = logging.getLogger(__name__)
type _KnowledgeFSRequestForwarder = Callable[
[str | None, str | None, bytes | None, bytes | None],
KnowledgeFSUpstreamResponse,
]
_MAX_PROXY_BODY_BYTES = 64 * 1024 * 1024
_RESPONSE_HEADER_ALLOWLIST = (
"Cache-Control",
@@ -138,25 +128,27 @@ def _translate_proxy_error(exc: Exception, *, tenant_id: str) -> NoReturn:
def _knowledge_fs_operation_access_required(
view: Callable[[KnowledgeFSAuthorization], ResponseReturnValue],
view: Callable[[KnowledgeFSMethod, str], ResponseReturnValue],
) -> Callable[[KnowledgeFSMethod, str], ResponseReturnValue]:
"""Authorize one declared operation before billing and request-body work."""
@wraps(view)
def decorated(method: KnowledgeFSMethod, upstream_path: str) -> ResponseReturnValue:
current_user, tenant_id = current_account_with_tenant()
try:
authorization = authorize_knowledge_fs_request(
account=current_user,
tenant_id=tenant_id,
method=method,
path=upstream_path,
)
operation = get_knowledge_fs_operation(method, upstream_path)
except KnowledgeFSRouteNotAllowedError as exc:
raise NotFound() from exc
current_user, tenant_id = current_account_with_tenant()
try:
authorize_knowledge_fs_request(
account=current_user,
tenant_id=tenant_id,
operation=operation,
)
except KnowledgeFSAccessDeniedError as exc:
_translate_proxy_error(exc, tenant_id=tenant_id)
return view(authorization)
return view(method, upstream_path)
return decorated
@@ -198,26 +190,21 @@ def _proxy_response(
"""Expose raw content, status, and allowlisted headers from KnowledgeFS.
Raises:
HTTPException: KnowledgeFS returns a status normalized by the operation contract.
BadGateway: KnowledgeFS rejects the configured server credential.
Forbidden: KnowledgeFS denies the account access to the requested resource.
"""
upstream = upstream_result.response
mapped_status = dict(upstream_result.operation.error_status_map).get(upstream.status_code)
if mapped_status is not None:
if upstream.status_code == HTTPStatus.UNAUTHORIZED:
upstream.close()
description = "KnowledgeFS upstream request failed"
if upstream.status_code == HTTPStatus.UNAUTHORIZED:
description = "KnowledgeFS authentication failed"
logger.error(
"KnowledgeFS rejected the Dify server credential with HTTP %s for tenant_id=%s",
upstream.status_code,
tenant_id,
)
exception_type = default_exceptions.get(mapped_status)
if exception_type is None:
exception = HTTPException(description)
exception.code = mapped_status
raise exception
raise exception_type(description)
logger.error(
"KnowledgeFS rejected the Dify server credential with HTTP %s for tenant_id=%s",
upstream.status_code,
tenant_id,
)
raise BadGateway("KnowledgeFS authentication failed")
if upstream.status_code == HTTPStatus.FORBIDDEN:
upstream.close()
raise Forbidden()
allowed_header_names = dict.fromkeys(
name.lower() for name in (*_RESPONSE_HEADER_ALLOWLIST, *contract_response_headers)
@@ -250,21 +237,26 @@ def _proxy_response(
return Response(content, status=upstream.status_code, headers=headers)
def _proxy_current_request(
*,
method: KnowledgeFSMethod,
tenant_id: str,
forward: _KnowledgeFSRequestForwarder,
) -> Response:
"""Forward the current raw request through one preconfigured service entry."""
def _proxy_request(method: KnowledgeFSMethod, upstream_path: str) -> Response:
"""Forward the current raw request and return its filtered upstream response.
The call performs one outbound KnowledgeFS request. Integration failures are
converted to Console HTTP exceptions for the outer JSON error adapter.
"""
if not dify_config.KNOWLEDGE_FS_ENABLED:
raise NotFound()
current_user, tenant_id = current_account_with_tenant()
try:
proxy_result = forward(
request.headers.get("Accept"),
request.content_type,
request.query_string or None,
_request_body() if method != "GET" else None,
proxy_result = proxy_knowledge_fs_request(
account=current_user,
method=method,
path=upstream_path,
tenant_id=tenant_id,
accept=request.headers.get("Accept"),
content_type=request.content_type,
query=request.query_string or None,
body=_request_body() if method != "GET" else None,
request_headers=request.headers,
)
except (
KnowledgeFSConfigurationError,
@@ -282,99 +274,20 @@ def _proxy_current_request(
)
def _proxy_request(
method: KnowledgeFSMethod,
upstream_path: str,
) -> Response:
"""Authorize and forward the current request through the combined service use case."""
if not dify_config.KNOWLEDGE_FS_ENABLED:
raise NotFound()
current_user, tenant_id = current_account_with_tenant()
def forward(
accept: str | None,
content_type: str | None,
query: bytes | None,
body: bytes | None,
) -> KnowledgeFSUpstreamResponse:
return proxy_knowledge_fs_request(
account=current_user,
method=method,
path=upstream_path,
tenant_id=tenant_id,
accept=accept,
content_type=content_type,
query=query,
body=body,
request_headers=request.headers,
)
return _proxy_current_request(method=method, tenant_id=tenant_id, forward=forward)
def _proxy_authorized_request(authorization: KnowledgeFSAuthorization) -> Response:
"""Forward the current request using one previously authorized operation capability.
Args:
authorization: Request-scoped capability produced before billing and body parsing.
Returns:
The filtered response returned by KnowledgeFS.
Raises:
HTTPException: The integration is disabled or forwarding fails.
"""
operation = authorization.operation
tenant_id = authorization.tenant_id
def forward(
accept: str | None,
content_type: str | None,
query: bytes | None,
body: bytes | None,
) -> KnowledgeFSUpstreamResponse:
return proxy_authorized_knowledge_fs_request(
authorization=authorization,
accept=accept,
content_type=content_type,
query=query,
body=body,
request_headers=request.headers,
)
return _proxy_current_request(method=operation.method, tenant_id=tenant_id, forward=forward)
@_knowledge_fs_enabled
@_knowledge_fs_operation_access_required
@cloud_edition_billing_rate_limit_check("knowledge")
def _proxy_knowledge_fs_non_get(
authorization: KnowledgeFSAuthorization,
method: KnowledgeFSMethod,
upstream_path: str,
) -> ResponseReturnValue:
"""Apply knowledge billing checks to one allowlisted non-GET operation."""
return _proxy_authorized_request(authorization)
return _proxy_request(method, upstream_path)
@bp.route(
"/knowledge-fs/<path:upstream_path>",
methods=["OPTIONS"],
provide_automatic_options=False,
)
@_console_api_errors
@_knowledge_fs_enabled
def proxy_knowledge_fs_options(upstream_path: str) -> ResponseReturnValue:
"""Complete a CORS preflight only for an enabled Console operation."""
requested_method = cast(KnowledgeFSMethod, request.headers.get("Access-Control-Request-Method", "").upper())
try:
get_knowledge_fs_operation(requested_method, upstream_path)
except KnowledgeFSRouteNotAllowedError as exc:
raise NotFound() from exc
return Response(status=HTTPStatus.NO_CONTENT)
@bp.route(
"/knowledge-fs/<path:upstream_path>",
methods=["GET"],
methods=["GET", "OPTIONS"],
provide_automatic_options=False,
)
@_console_api_errors
+27 -10
View File
@@ -4,16 +4,20 @@ from http import HTTPStatus
from flask import redirect
from flask_restx import Resource
from pydantic import BaseModel, Field
from werkzeug.exceptions import Conflict, Forbidden, NotFound
from werkzeug.exceptions import Conflict, NotFound
from controllers.common.fields import RedirectResponse
from controllers.common.schema import register_response_schema_models, register_schema_models
from controllers.console import console_ns
from controllers.console.wraps import (
RBACPermission,
RBACResourceScope,
account_initialization_required,
cloud_edition_billing_enabled,
cloud_edition_billing_paid_plan_required,
is_admin_or_owner_required,
only_edition_cloud,
rbac_permission_required,
setup_required,
)
from extensions.ext_database import db
@@ -21,7 +25,6 @@ from fields.base import ResponseModel
from libs.archive_storage import get_export_storage
from libs.helper import dump_response
from libs.login import current_account_with_tenant, login_required
from models import TenantAccountRole
from services.retention.workflow_run.archive_download_preparation import ARCHIVE_DOWNLOAD_MIME_TYPE
from services.retention.workflow_run.archive_download_task_cache import (
WorkflowRunArchiveDownloadStatus,
@@ -95,13 +98,11 @@ register_response_schema_models(
)
def _current_owner_or_admin_ids() -> tuple[str, str]:
"""Return current Cloud workspace IDs for an owner or admin, independently of enterprise RBAC."""
def _current_ids() -> tuple[str, str]:
"""Return current `(tenant_id, account_id)` or raise when no workspace is selected."""
current_user, current_tenant_id = current_account_with_tenant()
if not current_tenant_id:
raise NotFound("Current workspace not found")
if not TenantAccountRole.is_privileged_role(current_user.current_role):
raise Forbidden()
return current_tenant_id, current_user.id
@@ -123,8 +124,12 @@ class WorkflowRunArchivesApi(Resource):
@only_edition_cloud
@cloud_edition_billing_enabled
@cloud_edition_billing_paid_plan_required
@is_admin_or_owner_required
@rbac_permission_required(
RBACResourceScope.WORKSPACE, RBACPermission.WORKSPACE_ROLE_MANAGE, resource_required=False
)
def get(self):
tenant_id, _ = _current_owner_or_admin_ids()
tenant_id, _ = _current_ids()
return dump_response(WorkflowRunArchiveListResponse, list_workflow_run_archives(db.session(), tenant_id))
@@ -144,8 +149,12 @@ class WorkflowRunArchiveDownloadsApi(Resource):
@only_edition_cloud
@cloud_edition_billing_enabled
@cloud_edition_billing_paid_plan_required
@is_admin_or_owner_required
@rbac_permission_required(
RBACResourceScope.WORKSPACE, RBACPermission.WORKSPACE_ROLE_MANAGE, resource_required=False
)
def post(self):
tenant_id, account_id = _current_owner_or_admin_ids()
tenant_id, account_id = _current_ids()
payload = WorkflowRunArchiveDownloadPayload.model_validate(console_ns.payload or {})
try:
task = create_workflow_run_archive_download_task(
@@ -171,8 +180,12 @@ class WorkflowRunArchiveDownloadApi(Resource):
@only_edition_cloud
@cloud_edition_billing_enabled
@cloud_edition_billing_paid_plan_required
@is_admin_or_owner_required
@rbac_permission_required(
RBACResourceScope.WORKSPACE, RBACPermission.WORKSPACE_ROLE_MANAGE, resource_required=False
)
def get(self, download_id: str):
tenant_id, _ = _current_owner_or_admin_ids()
tenant_id, _ = _current_ids()
try:
task = get_workflow_run_archive_download_task(tenant_id=tenant_id, download_id=download_id)
except WorkflowRunArchiveDownloadTaskNotFoundError as exc:
@@ -196,8 +209,12 @@ class WorkflowRunArchiveDownloadFileApi(Resource):
@only_edition_cloud
@cloud_edition_billing_enabled
@cloud_edition_billing_paid_plan_required
@is_admin_or_owner_required
@rbac_permission_required(
RBACResourceScope.WORKSPACE, RBACPermission.WORKSPACE_ROLE_MANAGE, resource_required=False
)
def get(self, download_id: str):
tenant_id, _ = _current_owner_or_admin_ids()
tenant_id, _ = _current_ids()
try:
task = get_ready_workflow_run_archive_download_task(tenant_id=tenant_id, download_id=download_id)
except WorkflowRunArchiveDownloadTaskNotFoundError as exc:
@@ -67,15 +67,6 @@ from services.plugin.plugin_parameter_service import PluginParameterService
from services.plugin.plugin_permission_service import PluginPermissionService
from services.tools.tools_transform_service import ToolTransformService
_PLUGIN_PACKAGE_UPLOAD_PARAMS = {
"pkg": {
"description": "Plugin package to upload",
"in": "formData",
"type": "file",
"required": True,
}
}
class AutoUpgradeSettingsResponse(TypedDict):
strategy_setting: TenantPluginAutoUpgradeStrategySetting
@@ -654,7 +645,6 @@ class PluginAssetApi(Resource):
@console_ns.route("/workspaces/current/plugin/upload/pkg")
class PluginUploadFromPkgApi(Resource):
@console_ns.doc(consumes=["multipart/form-data"], params=_PLUGIN_PACKAGE_UPLOAD_PARAMS)
@console_ns.response(200, "Success", console_ns.models[PluginDecodeResponse.__name__])
@setup_required
@login_required
+7 -17
View File
@@ -1,6 +1,6 @@
from typing import Any, Self
from pydantic import AliasChoices, Field
from pydantic import AliasChoices, Field, computed_field
from sqlalchemy import select
from werkzeug.exceptions import Forbidden
@@ -9,13 +9,11 @@ from controllers.common.schema import register_response_schema_models
from controllers.web import web_ns
from controllers.web.wraps import WebApiResource
from extensions.ext_database import db
from extensions.storage.storage_type import StorageType
from fields.base import ResponseModel
from libs.helper import build_icon_url
from models.account import Tenant, TenantStatus
from models.model import App, EndUser, IconType, Site
from models.model import App, EndUser, Site
from services.feature_service import FeatureModel, FeatureService
from services.file_service import FileService
class WebSiteResponse(ResponseModel):
@@ -34,7 +32,11 @@ class WebSiteResponse(ResponseModel):
prompt_public: bool | None = None
show_workflow_steps: bool | None = None
use_icon_as_answer_icon: bool | None = None
icon_url: str | None = None
@computed_field(return_type=str | None) # type: ignore[prop-decorator]
@property
def icon_url(self) -> str | None:
return build_icon_url(self.icon_type, self.icon)
class WebModelConfigResponse(ResponseModel):
@@ -86,7 +88,6 @@ class WebAppSiteResponse(ResponseModel):
end_user_id: str | None,
features: FeatureModel,
can_replace_logo: bool,
icon_url: str | None = None,
) -> Self:
custom_config = None
if can_replace_logo:
@@ -101,7 +102,6 @@ class WebAppSiteResponse(ResponseModel):
)
site_response = WebSiteResponse.model_validate(site, from_attributes=True)
site_response.icon_url = icon_url if icon_url is not None else build_icon_url(site.icon_type, site.icon)
if features.billing.enabled and not features.webapp_copyright_enabled:
site_response.copyright = None
site_response.input_placeholder = None
@@ -123,15 +123,6 @@ register_response_schema_models(
)
def _build_site_icon_url(*, site: Site, tenant_id: str) -> str | None:
"""Use direct S3 URLs only in Cloud Mode and preserve preview URLs elsewhere."""
if site.icon_type != IconType.IMAGE or not site.icon:
return None
if dify_config.EDITION == "CLOUD" and StorageType(dify_config.STORAGE_TYPE) == StorageType.S3:
return FileService(db.engine).get_file_presigned_url(file_id=site.icon, tenant_id=tenant_id)
return build_icon_url(site.icon_type, site.icon)
@web_ns.route("/site")
class AppSiteApi(WebApiResource):
@web_ns.doc("Get App Site Info")
@@ -168,5 +159,4 @@ class AppSiteApi(WebApiResource):
end_user_id=end_user.id,
features=features,
can_replace_logo=features.can_replace_logo,
icon_url=_build_site_icon_url(site=site, tenant_id=tenant.id),
).model_dump(mode="json")
+29 -14
View File
@@ -682,27 +682,42 @@ class AgentAppGenerator(MessageBasedAppGenerator):
if draft_type == AgentConfigDraftType.DEBUG_BUILD.value
else AgentConfigDraftType.DRAFT
)
if effective_draft_type == AgentConfigDraftType.DRAFT:
from services.agent.composer_service import AgentComposerService
return AgentComposerService.get_or_create_normal_agent_draft(
session=session,
tenant_id=tenant_id,
agent=agent,
created_by=agent.updated_by or agent.created_by,
)
if not account_id:
raise AgentAppGeneratorError("Build draft requires an account user")
stmt = select(AgentConfigDraft).where(
AgentConfigDraft.tenant_id == tenant_id,
AgentConfigDraft.agent_id == agent.id,
AgentConfigDraft.draft_type == AgentConfigDraftType.DEBUG_BUILD,
AgentConfigDraft.account_id == account_id,
AgentConfigDraft.draft_type == effective_draft_type,
)
if effective_draft_type == AgentConfigDraftType.DEBUG_BUILD:
if not account_id:
raise AgentAppGeneratorError("Build draft requires an account user")
stmt = stmt.where(AgentConfigDraft.account_id == account_id)
else:
stmt = stmt.where(AgentConfigDraft.account_id.is_(None))
draft = session.scalar(stmt.order_by(AgentConfigDraft.updated_at.desc()).limit(1))
if draft is not None:
return draft
raise AgentAppGeneratorError("Agent build draft not found")
if effective_draft_type == AgentConfigDraftType.DEBUG_BUILD:
raise AgentAppGeneratorError("Agent build draft not found")
_, snapshot, agent_soul = AgentAppGenerator._resolve_agent_by_id(
tenant_id=tenant_id,
agent_id=agent.id,
snapshot_id=agent.active_config_snapshot_id,
session=session,
)
draft = AgentConfigDraft(
tenant_id=tenant_id,
agent_id=agent.id,
draft_type=AgentConfigDraftType.DRAFT,
account_id=None,
draft_owner_key="",
base_snapshot_id=snapshot.id,
config_snapshot=agent_soul,
created_by=agent.created_by,
updated_by=agent.updated_by,
)
session.add(draft)
session.flush()
return draft
@staticmethod
def _resolve_agent_by_id(
+1 -14
View File
@@ -1,7 +1,7 @@
import inspect
import json
import logging
from collections.abc import Callable, Generator, Mapping
from collections.abc import Callable, Generator
from typing import Any, cast
from urllib.parse import unquote
@@ -23,7 +23,6 @@ from core.plugin.impl.exc import (
PluginLLMPollingUnsupportedError,
PluginNotFoundError,
PluginPermissionDeniedError,
PluginRuntimeError,
PluginUniqueIdentifierError,
)
from core.trigger.errors import (
@@ -376,18 +375,6 @@ class BasePluginClient:
# type `PluginLLMPollingUnsupportedError`.
case PluginLLMPollingUnsupportedError.__name__:
raise PluginLLMPollingUnsupportedError(description=error_object.get("message"))
case PluginRuntimeError.__name__:
args = error_object.get("args")
lambda_request_id = args.get("request_id") if isinstance(args, Mapping) else None
if not isinstance(lambda_request_id, str):
lambda_request_id = None
runtime_message = error_object.get("message")
if not isinstance(runtime_message, str):
runtime_message = "Plugin runtime request failed"
raise PluginRuntimeError(
description=runtime_message,
lambda_request_id=lambda_request_id,
)
case _:
raise PluginInvokeError(description=message)
case PluginDaemonInternalServerError.__name__:
-12
View File
@@ -49,18 +49,6 @@ class PluginDaemonBadRequestError(PluginDaemonClientSideError):
description: str = "Bad Request"
class PluginRuntimeError(PluginDaemonInternalError):
"""A plugin runtime failed before it could return a valid plugin response."""
lambda_request_id: str | None
def __init__(self, description: str, lambda_request_id: str | None = None) -> None:
self.lambda_request_id = lambda_request_id
if lambda_request_id:
description = description.replace(f"RequestId: {lambda_request_id} Error: ", "", 1)
super().__init__(description)
class PluginInvokeError(PluginDaemonClientSideError, ValueError):
description: str = "Invoke Error"
+1 -4
View File
@@ -131,10 +131,7 @@ class WaterCrawlAPIClient(BaseAPIClient):
content_type = response.headers.get("Content-Type", "")
media_type = content_type.split(";", 1)[0].strip().lower()
if media_type == "application/json":
try:
return response.json() or {}
except ValueError as exc:
raise ValueError("Invalid JSON response from WaterCrawl") from exc
return response.json() or {}
if media_type == "application/octet-stream":
return response.content
@@ -1,15 +1,5 @@
"""WaterCrawl domain exceptions.
These exceptions are constructed from upstream HTTP responses, which may be
JSON API errors or plain text/HTML proxy errors. Keep the exception type stable
even when the body is not JSON so callers can handle WaterCrawl failures by
domain type instead of low-level parser errors.
"""
import json
from typing import Any, override
from httpx import Response
from typing import override
class WaterCrawlError(Exception):
@@ -17,16 +7,11 @@ class WaterCrawlError(Exception):
class WaterCrawlBadRequestError(WaterCrawlError):
def __init__(self, response: Response):
def __init__(self, response):
self.status_code = response.status_code
self.response = response
try:
data: Any = response.json()
except ValueError:
data = {}
if not isinstance(data, dict):
data = {}
self.message = data.get("message") or response.text or "Unknown error occurred"
data = response.json()
self.message = data.get("message", "Unknown error occurred")
self.errors = data.get("errors", {})
super().__init__(self.message)
-1
View File
@@ -61,7 +61,6 @@ class RBACPermission(StrEnum):
WORKSPACE_ROLE_MANAGE = "workspace_role_manage"
API_EXTENSION_MANAGE = "api_extension_manage"
CUSTOMIZATION_MANAGE = "customization_manage"
AGENT_MANAGE = "agent_manage"
SNIPPETS_CREATE_AND_MODIFY = "snippets_create_and_modify"
SNIPPETS_MANAGE = "snippets_management"
+15 -28
View File
@@ -58,6 +58,7 @@ class Tool(ABC):
if self.runtime and self.runtime.runtime_parameters:
tool_parameters.update(self.runtime.runtime_parameters)
# try parse tool parameters into the correct type
tool_parameters = self._transform_tool_parameters_type(tool_parameters)
result = self._invoke(
@@ -86,14 +87,14 @@ class Tool(ABC):
return result
def _transform_tool_parameters_type(self, tool_parameters: dict[str, Any]) -> dict[str, Any]:
"""Transform declared tool parameter values without resolving runtime schemas."""
"""
Transform tool parameters type
"""
# Temp fix for the issue that the tool parameters will be converted to empty while validating the credentials
result = deepcopy(tool_parameters)
for parameter in self.entity.parameters or []:
if parameter.name in tool_parameters:
if parameter.multiple:
result[parameter.name] = parameter.init_frontend_parameter(result.get(parameter.name))
else:
result[parameter.name] = parameter.type.cast_value(tool_parameters[parameter.name])
result[parameter.name] = parameter.type.cast_value(tool_parameters[parameter.name])
return result
@@ -195,31 +196,17 @@ class Tool(ABC):
}:
continue
is_multiple_select = parameter.multiple and parameter.type in {
ToolParameter.ToolParameterType.SELECT,
ToolParameter.ToolParameterType.DYNAMIC_SELECT,
}
if is_multiple_select:
item_schema: dict[str, Any] = {"type": "string"}
if parameter.type == ToolParameter.ToolParameterType.SELECT and parameter.options:
item_schema["enum"] = [option.value for option in parameter.options]
parameter_schema: dict[str, Any] = {"type": "array", "items": item_schema}
else:
parameter_schema = (
{
"type": parameter.type.as_normal_type(),
"description": parameter.llm_description or "",
}
if parameter.input_schema is None
else deepcopy(parameter.input_schema)
)
parameter_schema: dict[str, Any] = (
{
"type": parameter.type.as_normal_type(),
"description": parameter.llm_description or "",
}
if parameter.input_schema is None
else deepcopy(parameter.input_schema)
)
parameter_schema.setdefault("description", parameter.llm_description or "")
if (
not is_multiple_select
and parameter.type == ToolParameter.ToolParameterType.SELECT
and parameter.options
):
if parameter.type == ToolParameter.ToolParameterType.SELECT and parameter.options:
parameter_schema["enum"] = [option.value for option in parameter.options]
schema["properties"][parameter.name] = parameter_schema
+5 -36
View File
@@ -292,7 +292,9 @@ class ToolInvokeMessageBinary(BaseModel):
class ToolParameter(PluginParameter):
"""Tool-specific parameter declaration and invocation-value normalization."""
"""
Overrides type
"""
class ToolParameterType(StrEnum):
"""
@@ -331,28 +333,12 @@ class ToolParameter(PluginParameter):
LLM = auto() # will be set by LLM
type: ToolParameterType = Field(..., description="The type of the parameter")
multiple: bool = Field(
default=False,
description="Whether the parameter is multiple select, only valid for select or dynamic-select type",
)
human_description: I18nObject | None = Field(default=None, description="The description presented to the user")
form: ToolParameterForm = Field(..., description="The form of the parameter, schema/form/llm")
llm_description: str | None = None
# MCP object and array type parameters use this field to store the schema
input_schema: dict[str, Any] | None = None
@model_validator(mode="after")
def validate_multiple(self) -> ToolParameter:
supports_multiple = self.type in {
self.ToolParameterType.SELECT,
self.ToolParameterType.DYNAMIC_SELECT,
}
if self.multiple and not supports_multiple:
raise ValueError("multiple is only valid for select and dynamic-select parameters")
if supports_multiple and self.default is not None and (isinstance(self.default, list) != self.multiple):
raise ValueError("default must be a list exactly when multiple is true")
return self
@classmethod
def get_simple_instance(
cls,
@@ -392,25 +378,8 @@ class ToolParameter(PluginParameter):
options=option_objs,
)
def init_frontend_parameter(self, value: Any) -> Any:
"""Normalize a value against this tool parameter's full declaration."""
if not self.multiple:
return init_frontend_parameter(self, self.type, value)
parameter_value = self.default if value is None else value
if parameter_value is None:
parameter_value = []
if not isinstance(parameter_value, list):
raise ValueError(f"tool parameter {self.name} must be a list when multiple is true")
if not all(isinstance(item, str) for item in parameter_value):
raise ValueError(f"tool parameter {self.name} must contain only strings")
if self.required and not parameter_value:
raise ValueError(f"tool parameter {self.name} not found in tool config")
if self.type == self.ToolParameterType.SELECT:
options = [option.value for option in self.options]
if any(item not in options for item in parameter_value):
raise ValueError(f"tool parameter {self.name} value {parameter_value} not in options {options}")
return parameter_value
def init_frontend_parameter(self, value: Any):
return init_frontend_parameter(self, self.type, value)
class ToolProviderIdentity(BaseModel):
@@ -2,7 +2,7 @@ from __future__ import annotations
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Any, Final, Literal, Protocol
from typing import Any, Literal, Protocol, cast
from dify_agent.layers.dify_core_tools import DifyCoreToolConfig, DifyCoreToolProviderType, DifyCoreToolsLayerConfig
from dify_agent.layers.dify_plugin import (
@@ -29,13 +29,6 @@ from models.provider_ids import ToolProviderID
from models.tools import WorkflowToolProvider
from services.tools.mcp_tools_manage_service import MCPToolManageService
_CORE_TOOL_PROVIDER_TYPES: Final[dict[ToolProviderType, DifyCoreToolProviderType]] = {
ToolProviderType.BUILT_IN: "builtin",
ToolProviderType.API: "api",
ToolProviderType.WORKFLOW: "workflow",
ToolProviderType.MCP: "mcp",
}
class WorkflowAgentDifyToolsBuildError(ValueError):
"""Raised when Agent Soul tools cannot be prepared for Agent backend."""
@@ -242,7 +235,7 @@ class WorkflowAgentDifyToolsBuilder:
if tool_config.tool_name is not None:
expanded.append(tool_config)
continue
provider_type = tool_config.provider_type
provider_type = ToolProviderType.value_of(tool_config.provider_type)
provider_id = self._provider_id(tool_config)
try:
tool_names = self._provider_declared_tool_names(
@@ -286,7 +279,7 @@ class WorkflowAgentDifyToolsBuilder:
tenant_id: str,
tool_config: AgentSoulDifyToolConfig,
) -> AgentSoulDifyToolConfig:
if tool_config.provider_type is not ToolProviderType.MCP:
if tool_config.provider_type != ToolProviderType.MCP.value:
return tool_config
provider_id = self._mcp_provider_id_resolver(tenant_id=tenant_id, provider_id=self._provider_id(tool_config))
return tool_config.model_copy(update={"provider_id": provider_id, "plugin_id": None, "provider": None})
@@ -333,7 +326,7 @@ class WorkflowAgentDifyToolsBuilder:
def _to_agent_tool_entity(tool_config: AgentSoulDifyToolConfig) -> AgentToolEntity:
assert tool_config.tool_name is not None
return AgentToolEntity(
provider_type=tool_config.provider_type,
provider_type=ToolProviderType.value_of(tool_config.provider_type),
provider_id=WorkflowAgentDifyToolsBuilder._provider_id(tool_config),
tool_name=tool_config.tool_name,
tool_parameters=dict(tool_config.runtime_parameters),
@@ -350,16 +343,24 @@ class WorkflowAgentDifyToolsBuilder:
@staticmethod
def _provider_key(tool_config: AgentSoulDifyToolConfig) -> tuple[ToolProviderType, str]:
return (tool_config.provider_type, WorkflowAgentDifyToolsBuilder._provider_id(tool_config))
return (
ToolProviderType.value_of(tool_config.provider_type),
WorkflowAgentDifyToolsBuilder._provider_id(tool_config),
)
@staticmethod
def _tool_layer_destination(tool_config: AgentSoulDifyToolConfig) -> Literal["plugin", "core"]:
provider_type = tool_config.provider_type
provider_type = ToolProviderType.value_of(tool_config.provider_type)
if provider_type is ToolProviderType.PLUGIN or (
provider_type is ToolProviderType.BUILT_IN and _is_plugin_provider_id(tool_config.provider_id)
):
return "plugin"
if provider_type in _CORE_TOOL_PROVIDER_TYPES:
if provider_type in {
ToolProviderType.BUILT_IN,
ToolProviderType.API,
ToolProviderType.WORKFLOW,
ToolProviderType.MCP,
}:
return "core"
if provider_type is ToolProviderType.DATASET_RETRIEVAL:
raise WorkflowAgentDifyToolsBuildError(
@@ -416,7 +417,7 @@ class WorkflowAgentDifyToolsBuilder:
) -> DifyCoreToolConfig:
parameters = self._prepared_parameters(tool_runtime)
return DifyCoreToolConfig(
provider_type=_CORE_TOOL_PROVIDER_TYPES[tool_config.provider_type],
provider_type=cast(DifyCoreToolProviderType, tool_config.provider_type),
provider_id=self._provider_id(tool_config),
tool_name=tool_config.tool_name or exposed_name,
credential_id=tool_config.credential_ref.id if tool_config.credential_ref else None,
+4 -134
View File
@@ -12,7 +12,6 @@ import json
import subprocess
import sys
import tempfile
from copy import deepcopy
from pathlib import Path
from typing import Any, Literal, TypedDict
@@ -25,16 +24,6 @@ LOCK_PATH = API_ROOT / "knowledge-fs-contract.lock.json"
DEFAULT_REPOSITORY = WORKSPACE_ROOT.parent / "knowledge-fs"
OPENAPI_METHODS = ("delete", "get", "head", "options", "patch", "post", "put", "trace")
PROXY_METHODS = frozenset({"delete", "get", "patch", "post", "put"})
CONSOLE_PROXY_ERROR_SCHEMA_NAME = "ConsoleProxyError"
CONSOLE_PROXY_ERROR_SCHEMA: dict[str, Any] = {
"type": "object",
"required": ["code", "message", "status"],
"properties": {
"code": {"type": "string"},
"message": {"type": "string"},
"status": {"type": "integer"},
},
}
class ContractDeclaration(TypedDict):
@@ -49,7 +38,6 @@ class ContractDeclaration(TypedDict):
request_headers: tuple[str, ...]
response_headers: tuple[str, ...]
response_media_types: tuple[str, ...]
error_status_map: tuple[tuple[int, int], ...]
type DeclarationField = Literal[
@@ -82,7 +70,6 @@ def main() -> None:
mode.add_argument("--check", action="store_true")
mode.add_argument("--update-lock", action="store_true")
parser.add_argument("--repository", type=Path, default=DEFAULT_REPOSITORY)
parser.add_argument("--output-openapi", type=Path)
args = parser.parse_args()
repository = args.repository.resolve()
@@ -114,15 +101,7 @@ def main() -> None:
)
document: dict[str, Any] = json.loads(openapi_content)
declarations = console_contract_declarations()
validate_declarations(document, declarations)
if args.output_openapi:
filtered_document = filter_openapi_document(document, declarations)
filtered_document["x-dify-source-openapi-sha256"] = openapi_sha256
filtered_document["x-dify-console-declarations-sha256"] = contract_declarations_sha256(declarations)
args.output_openapi.parent.mkdir(parents=True, exist_ok=True)
args.output_openapi.write_text(json.dumps(filtered_document, indent=2) + "\n")
validate_declarations(document, console_contract_declarations())
if args.update_lock:
LOCK_PATH.write_text(
@@ -168,7 +147,8 @@ def validate_declarations(document: dict[str, Any], declarations: tuple[Contract
raise ValueError(f"KnowledgeFS OpenAPI path must be absolute: {path}")
if method not in PROXY_METHODS:
raise ValueError(f"KnowledgeFS proxy does not support {method.upper()} {path}")
expected: dict[DeclarationField, object] = {
expected: ContractDeclaration = {
"operation_id": operation_id,
"method": method.upper(),
"path": path[1:],
"required_scope": required_scope(operation),
@@ -186,114 +166,11 @@ def validate_declarations(document: dict[str, Any], declarations: tuple[Contract
f"KnowledgeFS operation {operation_id} field {field} drifted: "
f"expected {expected_value!r}, received {received_value!r}"
)
validate_error_status_map(operation_id, declaration["error_status_map"])
def filter_openapi_document(
document: dict[str, Any],
declarations: tuple[ContractDeclaration, ...],
) -> dict[str, Any]:
"""Return a code-generation document containing only Console-allowlisted operations."""
filtered_document: dict[str, Any] = {
key: value for key, value in document.items() if key not in {"components", "paths"}
}
source_paths = document.get("paths", {})
filtered_paths: dict[str, Any] = {}
for declaration in declarations:
path = f"/{declaration['path']}"
method = declaration["method"].lower()
source_path_item = source_paths[path]
path_metadata = {key: value for key, value in source_path_item.items() if key not in OPENAPI_METHODS}
filtered_path_item = filtered_paths.setdefault(path, path_metadata)
filtered_operation = deepcopy(source_path_item[method])
_rewrite_proxy_error_responses(filtered_operation, declaration["error_status_map"])
filtered_path_item[method] = filtered_operation
filtered_document["paths"] = filtered_paths
source_components = document.get("components", {})
filtered_components = {key: value for key, value in source_components.items() if key != "schemas"}
source_schemas = source_components.get("schemas", {})
available_schemas = {**source_schemas, CONSOLE_PROXY_ERROR_SCHEMA_NAME: CONSOLE_PROXY_ERROR_SCHEMA}
schema_names = _referenced_schema_names(filtered_paths, available_schemas)
filtered_components["schemas"] = {
name: schema for name, schema in available_schemas.items() if name in schema_names
}
filtered_document["components"] = filtered_components
return filtered_document
def validate_error_status_map(operation_id: str, error_status_map: tuple[tuple[int, int], ...]) -> None:
"""Validate the status normalization advertised by one Console operation."""
upstream_statuses: set[int] = set()
for upstream_status, console_status in error_status_map:
if upstream_status in upstream_statuses:
raise ValueError(f"KnowledgeFS operation {operation_id} has duplicate error status: {upstream_status}")
if not 400 <= upstream_status <= 599 or not 400 <= console_status <= 599:
raise ValueError(f"KnowledgeFS operation {operation_id} has invalid error status mapping")
upstream_statuses.add(upstream_status)
def _rewrite_proxy_error_responses(
operation: dict[str, Any],
error_status_map: tuple[tuple[int, int], ...],
) -> None:
responses = operation.setdefault("responses", {})
proxy_error_response = {
"description": "Error normalized by the Dify Console KnowledgeFS proxy.",
"content": {
"application/json": {"schema": {"$ref": f"#/components/schemas/{CONSOLE_PROXY_ERROR_SCHEMA_NAME}"}}
},
}
for upstream_status, console_status in error_status_map:
existing_target = responses.get(str(console_status)) if upstream_status != console_status else None
responses.pop(str(upstream_status), None)
normalized_response: dict[str, Any] = deepcopy(proxy_error_response)
existing_schema = (
existing_target.get("content", {}).get("application/json", {}).get("schema")
if isinstance(existing_target, dict)
else None
)
if existing_schema is not None:
normalized_response["content"]["application/json"]["schema"] = {
"oneOf": [
deepcopy(existing_schema),
{"$ref": f"#/components/schemas/{CONSOLE_PROXY_ERROR_SCHEMA_NAME}"},
]
}
responses[str(console_status)] = normalized_response
def _referenced_schema_names(value: Any, schemas: dict[str, Any]) -> set[str]:
reference_prefix = "#/components/schemas/"
selected: set[str] = set()
pending: list[Any] = [value]
while pending:
current = pending.pop()
if isinstance(current, list):
pending.extend(current)
continue
if not isinstance(current, dict):
continue
reference = current.get("$ref")
if isinstance(reference, str) and reference.startswith(reference_prefix):
name = reference.removeprefix(reference_prefix)
if name not in selected:
if name not in schemas:
raise ValueError(f"KnowledgeFS OpenAPI references missing schema: {name}")
selected.add(name)
pending.append(schemas[name])
pending.extend(current.values())
return selected
def console_contract_declarations() -> tuple[ContractDeclaration, ...]:
"""Return transport declarations from the runtime Console operation registry."""
from services.knowledge_fs_operations import KNOWLEDGE_FS_CONSOLE_OPERATIONS
from services.knowledge_fs_proxy import KNOWLEDGE_FS_CONSOLE_OPERATIONS
return tuple(
{
@@ -306,18 +183,11 @@ def console_contract_declarations() -> tuple[ContractDeclaration, ...]:
"request_headers": operation.request_headers,
"response_headers": operation.response_headers,
"response_media_types": operation.response_media_types,
"error_status_map": operation.error_status_map,
}
for operation in KNOWLEDGE_FS_CONSOLE_OPERATIONS
)
def contract_declarations_sha256(declarations: tuple[ContractDeclaration, ...]) -> str:
"""Return a stable digest for the runtime Console operation declarations."""
content = json.dumps(declarations, separators=(",", ":"), sort_keys=True).encode()
return sha256(content)
def response_kind(operation: dict[str, Any]) -> str:
media_types = response_media_types(operation)
if "text/event-stream" in media_types:
-1
View File
@@ -157,7 +157,6 @@ def init_app(app: DifyApp) -> Celery:
"tasks.regenerate_summary_index_task", # summary index regeneration
"tasks.initialize_created_app_rbac_access_task", # app access initialization
"tasks.install_default_plugins_task", # tenant default plugin installation
"tasks.refresh_billing_vector_space_task", # billing vector-space cache refresh
"tasks.app_generate.resume_agent_app_task", # ENG-635: Agent v2 chat ask_human resume
"tasks.workflow_run_archive_download_tasks", # workflow-run archive download preparation
]
-13
View File
@@ -119,19 +119,6 @@ class Storage:
def delete(self, filename: str):
return self.storage_runner.delete(filename)
def generate_presigned_url(
self,
filename: str,
*,
expires_in: int,
content_type: str | None = None,
) -> str:
return self.storage_runner.generate_presigned_url(
filename,
expires_in=expires_in,
content_type=content_type,
)
def scan(self, path: str, files: bool = True, directories: bool = False) -> list[str]:
return self.storage_runner.scan(path, files=files, directories=directories)
-18
View File
@@ -92,21 +92,3 @@ class AwsS3Storage(BaseStorage):
@override
def delete(self, filename: str):
self.client.delete_object(Bucket=self.bucket_name, Key=filename)
@override
def generate_presigned_url(
self,
filename: str,
*,
expires_in: int,
content_type: str | None = None,
) -> str:
params = {"Bucket": self.bucket_name, "Key": filename}
if content_type:
params["ResponseContentType"] = content_type
return self.client.generate_presigned_url(
"get_object",
Params=params,
ExpiresIn=expires_in,
)
-10
View File
@@ -31,16 +31,6 @@ class BaseStorage(ABC):
def delete(self, filename: str):
raise NotImplementedError
def generate_presigned_url(
self,
filename: str,
*,
expires_in: int,
content_type: str | None = None,
) -> str:
"""Generate a temporary direct-download URL when the backend supports it."""
raise NotImplementedError("This storage backend doesn't support presigned URLs")
def scan(self, path, files=True, directories=False) -> list[str]:
"""
Scan files and directories in the given path.
+2 -2
View File
@@ -1,5 +1,5 @@
{
"commit": "a0f50470612cc0b3656f89e4f2435aaf412e6e3b",
"openapiSha256": "f18910e9c45a64f0855e0643a7a626fb2889021b4f943458de86c6bd2469facb",
"commit": "4310e2d582d25e7de58183f27720afab01e123cf",
"openapiSha256": "5827ca930ce38462bfd1b2bef387efbf37eb7ffcaedde4558af2fbaeccbfbc4b",
"repository": "https://github.com/langgenius/knowledge-fs"
}
-17
View File
@@ -9,8 +9,6 @@ from werkzeug.http import HTTP_STATUS_CODES
from configs import dify_config
from core.errors.error import AppInvokeQuotaExceededError
from core.plugin.impl.exc import PluginRuntimeError
from extensions.ext_logging import get_request_id
from libs.flask_restx_compat import install_swagger_compatibility
from libs.token import build_force_logout_cookie_headers
@@ -102,20 +100,6 @@ def register_external_error_handlers(api: Api, body_formatter: ErrorBodyFormatte
data = {"code": "too_many_requests", "message": str(e), "status": status_code}
return _finalize(e, data, status_code), status_code
def handle_plugin_runtime_error(e: PluginRuntimeError):
got_request_exception.send(current_app, exception=e)
status_code = 502
details = {"request_id": get_request_id()}
if e.lambda_request_id:
details["lambda_request_id"] = e.lambda_request_id
data = {
"code": "plugin_runtime_error",
"message": e.description,
"details": details,
"status": status_code,
}
return _finalize(e, data, status_code), status_code
def handle_general_exception(e: Exception):
got_request_exception.send(current_app, exception=e)
@@ -137,7 +121,6 @@ def register_external_error_handlers(api: Api, body_formatter: ErrorBodyFormatte
api.errorhandler(HTTPException)(handle_http_exception)
api.errorhandler(ValueError)(handle_value_error)
api.errorhandler(AppInvokeQuotaExceededError)(handle_quota_exceeded)
api.errorhandler(PluginRuntimeError)(handle_plugin_runtime_error)
api.errorhandler(Exception)(handle_general_exception)
+2 -5
View File
@@ -221,11 +221,8 @@ def current_timestamp() -> int:
def email(email):
# Define a regex pattern for email addresses
pattern = r"^[\w\.!#$%&'*+\-/=?^_`{|}~]+@([\w-]+\.)+[\w-]{2,}$"
# Use re.fullmatch instead of re.match to reject trailing newlines.
# In Python, '$' matches at end-of-string OR just before a trailing newline,
# so re.match accepts "user@example.com\n". re.fullmatch requires the entire
# string to match, closing the mail header-injection vector. (#39234)
if re.fullmatch(pattern, email) is not None:
# Check if the email matches the pattern
if re.match(pattern, email) is not None:
return email
error = f"{email} is not a valid email."
+2 -2
View File
@@ -23,9 +23,9 @@ class PaginatedResult[T]:
@property
def pages(self) -> int:
if self.total == 0 or self.per_page == 0:
if self.per_page == 0:
return 0
return math.ceil(self.total / self.per_page)
return max(1, math.ceil(self.total / self.per_page))
@property
def has_next(self) -> bool:
@@ -1,77 +0,0 @@
"""scope agent debug conversations by draft type
Revision ID: d2825e7b9c10
Revises: b8c9d0e1f2a3
Create Date: 2026-07-22 15:00:00.000000
"""
import sqlalchemy as sa
from alembic import op
import models
# revision identifiers, used by Alembic.
revision = "d2825e7b9c10"
down_revision = "b8c9d0e1f2a3"
branch_labels = None
depends_on = None
def upgrade():
# Existing pointers have always represented Build chat because the Agent
# detail API exposes them as ``debug_conversation_id`` for that surface.
op.add_column(
"agent_debug_conversations",
sa.Column(
"draft_type",
sa.String(length=32),
nullable=False,
server_default=sa.text("'debug_build'"),
),
)
op.drop_constraint(
"agent_debug_conversation_agent_account_unique",
"agent_debug_conversations",
type_="unique",
)
op.create_unique_constraint(
"agent_debug_conversation_agent_account_draft_type_unique",
"agent_debug_conversations",
["tenant_id", "agent_id", "account_id", "draft_type"],
)
def downgrade():
debug_conversations = sa.table(
"agent_debug_conversations",
sa.column("tenant_id", models.types.StringUUID()),
sa.column("agent_id", models.types.StringUUID()),
sa.column("account_id", models.types.StringUUID()),
sa.column("draft_type", sa.String(length=32)),
)
build_conversations = debug_conversations.alias("build_conversations")
op.get_bind().execute(
sa.delete(debug_conversations).where(
debug_conversations.c.draft_type == "draft",
sa.exists(
sa.select(sa.literal(1)).where(
build_conversations.c.tenant_id == debug_conversations.c.tenant_id,
build_conversations.c.agent_id == debug_conversations.c.agent_id,
build_conversations.c.account_id == debug_conversations.c.account_id,
build_conversations.c.draft_type == "debug_build",
)
),
)
)
op.drop_constraint(
"agent_debug_conversation_agent_account_draft_type_unique",
"agent_debug_conversations",
type_="unique",
)
op.create_unique_constraint(
"agent_debug_conversation_agent_account_unique",
"agent_debug_conversations",
["tenant_id", "agent_id", "account_id"],
)
op.drop_column("agent_debug_conversations", "draft_type")
+3 -12
View File
@@ -222,13 +222,11 @@ class Agent(DefaultFieldsMixin, Base):
class AgentDebugConversation(DefaultFieldsMixin, Base):
"""Per-account, per-draft console debug conversation for an Agent App.
"""Per-account console debug conversation for an Agent App.
Agent App preview state must be isolated by editor account. The Agent row is
shared by everyone in the workspace, so this table owns the user-specific
conversation pointers used by console debug chat. ``draft`` is the Preview
conversation and ``debug_build`` is the Build conversation; they must never
share persisted messages or runtime sessions.
conversation pointer used by console debug chat.
"""
__tablename__ = "agent_debug_conversations"
@@ -238,8 +236,7 @@ class AgentDebugConversation(DefaultFieldsMixin, Base):
"tenant_id",
"agent_id",
"account_id",
"draft_type",
name="agent_debug_conversation_agent_account_draft_type_unique",
name="agent_debug_conversation_agent_account_unique",
),
Index("agent_debug_conversation_conversation_idx", "conversation_id"),
Index("agent_debug_conversation_account_idx", "tenant_id", "account_id"),
@@ -249,12 +246,6 @@ class AgentDebugConversation(DefaultFieldsMixin, Base):
agent_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
app_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
account_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
draft_type: Mapped[AgentConfigDraftType] = mapped_column(
EnumText(AgentConfigDraftType, length=32),
nullable=False,
default=AgentConfigDraftType.DEBUG_BUILD,
server_default=sa.text("'debug_build'"),
)
conversation_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
+15 -22
View File
@@ -7,7 +7,6 @@ from typing import Annotated, Any, Final, Literal, Self
from pydantic import BaseModel, ConfigDict, Field, WithJsonSchema, field_validator, model_validator
from core.rag.entities.metadata_entities import ConditionValue, SupportedComparisonOperator
from core.tools.entities.tool_entities import ToolProviderType
from core.workflow.file_reference import is_canonical_file_reference
from graphon.file import FileTransferMethod, FileType
@@ -44,28 +43,18 @@ _DECLARED_OUTPUT_CHILDREN_JSON_SCHEMA = {
},
"description": {"anyOf": [{"type": "string"}, {"type": "null"}]},
"required": {"type": "boolean"},
"file": {
"anyOf": [
{"type": "object", "additionalProperties": True},
{"type": "null"},
]
},
"file": {"type": "object", "additionalProperties": True},
"array_item": {
"anyOf": [
{
"type": "object",
"additionalProperties": True,
"properties": {
"type": {
"type": "string",
"enum": [item.value for item in DeclaredOutputType],
},
"description": {"anyOf": [{"type": "string"}, {"type": "null"}]},
"children": {"type": "array", "items": {"type": "object", "additionalProperties": True}},
},
"type": "object",
"additionalProperties": True,
"properties": {
"type": {
"type": "string",
"enum": [item.value for item in DeclaredOutputType],
},
{"type": "null"},
]
"description": {"anyOf": [{"type": "string"}, {"type": "null"}]},
"children": {"type": "array", "items": {"type": "object", "additionalProperties": True}},
},
},
"children": {"type": "array", "items": {"type": "object", "additionalProperties": True}},
},
@@ -664,7 +653,11 @@ class AgentSoulDifyToolConfig(BaseModel):
model_config = ConfigDict(extra="ignore")
enabled: bool = True
provider_type: ToolProviderType
# ``plugin`` remains the default for legacy Agent Soul payloads. The runtime
# now also accepts ``builtin`` / ``api`` / ``workflow`` / ``mcp`` here and
# routes them through ``dify.core.tools``; keeping the default narrow still
# makes a missing field resolve against the plugin provider table.
provider_type: str = "plugin"
provider_id: str | None = Field(default=None, max_length=255)
plugin_id: str | None = Field(default=None, max_length=255)
provider: str | None = Field(default=None, max_length=255)
+13 -31
View File
@@ -56,7 +56,6 @@ from .provider_ids import GenericProviderID
from .types import EnumText, LongText, StringUUID
if TYPE_CHECKING:
from .agent import Agent
from .workflow import Workflow
@@ -502,42 +501,25 @@ class App(Base):
Resolved via ``Agent.app_id`` so the console can open the Composer in
roster-detail mode from the app id. ``None`` for non-agent apps.
"""
agent = self.agent_app_binding_with_session(session=session)
return agent.id if agent else None
def agent_app_binding_with_session(self, *, session: Session, include_archived: bool = False) -> Agent | None:
"""For an Agent App (mode=agent), the Agent bound to it.
A roster Agent is bound through ``Agent.app_id``; a workflow-only Agent
is bound to its hidden runtime backing App through
``Agent.backing_app_id``. Callers branch on ``Agent.scope`` to tell the
public roster Agent App apart from the hidden backing App. Archived
Agents are excluded unless ``include_archived`` is set (authorization
gates must keep covering an Agent App after its Agent is archived).
``None`` for non-agent apps and unbound agent apps.
"""
if self.mode != AppMode.AGENT:
return None
from .agent import APP_BACKED_AGENT_SOURCES, Agent, AgentScope, AgentStatus
conditions = [
Agent.tenant_id == self.tenant_id,
sa.or_(
sa.and_(
Agent.app_id == self.id,
Agent.scope == AgentScope.ROSTER,
Agent.source.in_(APP_BACKED_AGENT_SOURCES),
),
sa.and_(
agent = session.scalar(
select(Agent).where(
Agent.tenant_id == self.tenant_id,
sa.or_(
sa.and_(
Agent.app_id == self.id,
Agent.scope == AgentScope.ROSTER,
Agent.source.in_(APP_BACKED_AGENT_SOURCES),
),
Agent.backing_app_id == self.id,
Agent.scope == AgentScope.WORKFLOW_ONLY,
),
),
]
if not include_archived:
conditions.append(Agent.status == AgentStatus.ACTIVE)
return session.scalar(select(Agent).where(*conditions).limit(1))
Agent.status == AgentStatus.ACTIVE,
)
)
return agent.id if agent else None
@property
def api_base_url(self) -> str:
+11 -74
View File
@@ -260,12 +260,7 @@ Get account avatar url
| 200 | Success | **application/json**: [AccountResponse](#accountresponse)<br> |
### [POST] /activate
**Accept an invitation without letting an existing session act for another account**
Activate account with invitation token
Token-only activation remains available for legacy clients. When the request already
carries a console session, that session must belong to the account encoded in the
invitation before the token is consumed or tenant membership is changed.
#### Request Body
@@ -536,7 +531,6 @@ Run a build-draft Agent App turn that asks the agent to push config updates
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 200 | Agent build draft | **application/json**: [AgentBuildDraftResponse](#agentbuilddraftresponse)<br> |
| 404 | Agent build draft not found | |
### [PUT] /agent/{agent_id}/build-draft
#### Parameters
@@ -964,12 +958,6 @@ Stop a running Agent App chat message generation
| ---- | ---------- | ----------- | -------- | ------ |
| agent_id | path | | Yes | string (uuid) |
#### Request Body
| Required | Schema |
| -------- | ------ |
| No | **application/json**: [AgentDebugConversationRefreshPayload](#agentdebugconversationrefreshpayload)<br> |
#### Responses
| Code | Description | Schema |
@@ -11344,12 +11332,6 @@ Returns permission flags that control workspace features like member invitations
| 200 | Success | **application/json**: [PluginDecodeResponse](#plugindecoderesponse)<br> |
### [POST] /workspaces/current/plugin/upload/pkg
#### Request Body
| Required | Schema |
| -------- | ------ |
| Yes | **multipart/form-data**: { **"pkg"**: binary }<br> |
#### Responses
| Code | Description | Schema |
@@ -13296,7 +13278,7 @@ Model class for AI model.
| maintainer | string | | No |
| max_active_requests | integer | | No |
| mode | string | | Yes |
| model_config | [AppModelConfigResponse](#appmodelconfigresponse) | | No |
| model_config | [ModelConfig](#modelconfig) | | No |
| name | string | | Yes |
| permission_keys | [ string ] | | No |
| role | string | | No |
@@ -13917,12 +13899,6 @@ Stable Agent Soul reference to one normalized skill archive.
| date | string | | Yes |
| message_count | integer | | Yes |
#### AgentDebugConversationRefreshPayload
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| draft_type | [AgentConfigDraftType](#agentconfigdrafttype) | Agent draft surface whose conversation should be refreshed | No |
#### AgentDebugConversationRefreshResponse
| Name | Type | Description | Required |
@@ -14761,7 +14737,7 @@ should send ``plugin_id`` + ``provider`` when available.
| plugin_id | string | | No |
| provider | string | | No |
| provider_id | string | | No |
| provider_type | [ToolProviderType](#toolprovidertype) | | Yes |
| provider_type | string, <br>**Default:** plugin | | No |
| runtime_parameters | object | | No |
| tool_name | string | | No |
@@ -15372,6 +15348,7 @@ This class is used to store the schema information of an api based tool.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| access_mode | string | | No |
| app_model_config | [ModelConfig](#modelconfig) | | No |
| created_at | integer | | No |
| created_by | string | | No |
| description | string | | No |
@@ -15381,8 +15358,7 @@ This class is used to store the schema information of an api based tool.
| icon_background | string | | No |
| id | string | | Yes |
| maintainer | string | | No |
| mode | string | | Yes |
| model_config | [AppModelConfigResponse](#appmodelconfigresponse) | | No |
| mode_compatible_with_agent | string | | Yes |
| name | string | | Yes |
| permission_keys | [ string ] | | No |
| tags | [ [Tag](#tag) ] | | No |
@@ -15444,7 +15420,7 @@ This class is used to store the schema information of an api based tool.
| maintainer | string | | No |
| max_active_requests | integer | | No |
| mode | string | | Yes |
| model_config | [AppModelConfigResponse](#appmodelconfigresponse) | | No |
| model_config | [ModelConfig](#modelconfig) | | No |
| name | string | | Yes |
| permission_keys | [ string ] | | No |
| site | [AppDetailSiteResponse](#appdetailsiteresponse) | | No |
@@ -15543,35 +15519,6 @@ AppMCPServer Status Enum
| ---- | ---- | ----------- | -------- |
| AppMCPServerStatus | string | AppMCPServer Status Enum | |
#### AppModelConfigResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| agent_mode | | | No |
| annotation_reply | | | No |
| chat_prompt_config | | | No |
| completion_prompt_config | | | No |
| created_at | integer | | No |
| created_by | string | | No |
| dataset_configs | | | No |
| dataset_query_variable | string | | No |
| external_data_tools | | | No |
| file_upload | | | No |
| model | | | No |
| more_like_this | | | No |
| opening_statement | string | | No |
| pre_prompt | string | | No |
| prompt_type | string | | No |
| retriever_resource | | | No |
| sensitive_word_avoidance | | | No |
| speech_to_text | | | No |
| suggested_questions | | | No |
| suggested_questions_after_answer | | | No |
| text_to_speech | | | No |
| updated_at | integer | | No |
| updated_by | string | | No |
| user_input_form | | | No |
#### AppNamePayload
| Name | Type | Description | Required |
@@ -17200,7 +17147,7 @@ about. Stage 4 §4.2.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| children | [ { **"array_item"**: , **"children"**: [ object ], **"description"**: , **"file"**: , **"name"**: string, **"required"**: boolean, **"type"**: string, <br>**Available values:** "array", "boolean", "file", "number", "object", "string" } ] | | No |
| children | [ { **"array_item"**: { **"children"**: [ object ], **"description"**: , **"type"**: string, <br>**Available values:** "array", "boolean", "file", "number", "object", "string" }, **"children"**: [ object ], **"description"**: , **"file"**: object, **"name"**: string, **"required"**: boolean, **"type"**: string, <br>**Available values:** "array", "boolean", "file", "number", "object", "string" } ] | | No |
| description | string | | No |
| type | [DeclaredOutputType](#declaredoutputtype) | | Yes |
@@ -17229,7 +17176,7 @@ code can call ``output.failure_strategy.on_failure`` without None-guards.
| ---- | ---- | ----------- | -------- |
| array_item | [DeclaredArrayItem](#declaredarrayitem) | | No |
| check | [DeclaredOutputCheckConfig](#declaredoutputcheckconfig) | | No |
| children | [ { **"array_item"**: , **"children"**: [ object ], **"description"**: , **"file"**: , **"name"**: string, **"required"**: boolean, **"type"**: string, <br>**Available values:** "array", "boolean", "file", "number", "object", "string" } ] | | No |
| children | [ { **"array_item"**: { **"children"**: [ object ], **"description"**: , **"type"**: string, <br>**Available values:** "array", "boolean", "file", "number", "object", "string" }, **"children"**: [ object ], **"description"**: , **"file"**: object, **"name"**: string, **"required"**: boolean, **"type"**: string, <br>**Available values:** "array", "boolean", "file", "number", "object", "string" } ] | | No |
| description | string | | No |
| failure_strategy | [DeclaredOutputFailureStrategy](#declaredoutputfailurestrategy) | | No |
| file | [DeclaredOutputFileConfig](#declaredoutputfileconfig) | | No |
@@ -17318,12 +17265,6 @@ Default model entity.
| tool_name | string | | Yes |
| type | string | | Yes |
#### DeploymentEdition
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| DeploymentEdition | string | | |
#### DismissNotificationPayload
| Name | Type | Description | Required |
@@ -22013,9 +21954,9 @@ The subscription constructor of the trigger provider
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| hash | string | | Yes |
| result | string | | Yes |
| updated_at | integer | | Yes |
| hash | string | | No |
| result | string | | No |
| updated_at | string | | No |
#### SystemConfigurationResponse
@@ -22032,7 +21973,6 @@ Model class for provider system configuration response.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| branding | [BrandingModel](#brandingmodel) | | Yes |
| deployment_edition | [DeploymentEdition](#deploymentedition) | | Yes |
| enable_app_deploy | boolean | | Yes |
| enable_change_email | boolean, <br>**Default:** true | | Yes |
| enable_collaboration_mode | boolean, <br>**Default:** true | | Yes |
@@ -22048,7 +21988,6 @@ Model class for provider system configuration response.
| is_allow_create_workspace | boolean | | Yes |
| is_allow_register | boolean | | Yes |
| is_email_setup | boolean | | Yes |
| knowledge_fs_enabled | boolean | | Yes |
| license | [LicenseModel](#licensemodel) | | Yes |
| max_plugin_package_size | integer, <br>**Default:** 15728640 | | Yes |
| plugin_installation_permission | [PluginInstallationPermissionModel](#plugininstallationpermissionmodel) | | Yes |
@@ -22342,7 +22281,7 @@ Tool label
#### ToolParameter
Tool-specific parameter declaration and invocation-value normalization.
Overrides type
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
@@ -22355,7 +22294,6 @@ Tool-specific parameter declaration and invocation-value normalization.
| llm_description | string | | No |
| max | number<br>integer | | No |
| min | number<br>integer | | No |
| multiple | boolean | Whether the parameter is multiple select, only valid for select or dynamic-select type | No |
| name | string | The name of the parameter | Yes |
| options | [ [PluginParameterOption](#pluginparameteroption) ] | | No |
| placeholder | [I18nObject](#i18nobject) | The placeholder presented to the user | No |
@@ -23002,7 +22940,6 @@ in form definiton, or a variable while the workflow is running.
| ---- | ---- | ----------- | -------- |
| allow_email_code_login | boolean | | Yes |
| allow_email_password_login | boolean | | Yes |
| allow_public_access | boolean, <br>**Default:** true | | Yes |
| allow_sso | boolean | | Yes |
| enabled | boolean | | Yes |
| sso_config | [WebAppAuthSSOModel](#webappauthssomodel) | | Yes |
+1 -10
View File
@@ -1058,12 +1058,6 @@ Button styles for user actions.
| auto_generate | boolean | Automatically generate the conversation name. When `true`, the `name` field is ignored. | No |
| name | string | Conversation name. Required when `auto_generate` is `false`. | No |
#### DeploymentEdition
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| DeploymentEdition | string | | |
#### EmailCodeLoginSendPayload
| Name | Type | Description | Required |
@@ -1573,7 +1567,6 @@ Default configuration for form inputs.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| branding | [BrandingModel](#brandingmodel) | | Yes |
| deployment_edition | [DeploymentEdition](#deploymentedition) | | Yes |
| enable_app_deploy | boolean | | Yes |
| enable_change_email | boolean, <br>**Default:** true | | Yes |
| enable_collaboration_mode | boolean, <br>**Default:** true | | Yes |
@@ -1589,7 +1582,6 @@ Default configuration for form inputs.
| is_allow_create_workspace | boolean | | Yes |
| is_allow_register | boolean | | Yes |
| is_email_setup | boolean | | Yes |
| knowledge_fs_enabled | boolean | | Yes |
| license | [LicenseModel](#licensemodel) | | Yes |
| max_plugin_package_size | integer, <br>**Default:** 15728640 | | Yes |
| plugin_installation_permission | [PluginInstallationPermissionModel](#plugininstallationpermissionmodel) | | Yes |
@@ -1651,7 +1643,6 @@ in form definiton, or a variable while the workflow is running.
| ---- | ---- | ----------- | -------- |
| allow_email_code_login | boolean | | Yes |
| allow_email_password_login | boolean | | Yes |
| allow_public_access | boolean, <br>**Default:** true | | Yes |
| allow_sso | boolean | | Yes |
| enabled | boolean | | Yes |
| sso_config | [WebAppAuthSSOModel](#webappauthssomodel) | | Yes |
@@ -1741,7 +1732,7 @@ in form definiton, or a variable while the workflow is running.
| icon | string | | No |
| icon_background | string | | No |
| icon_type | string | | No |
| icon_url | string | | No |
| icon_url | string | | Yes |
| input_placeholder | string | | No |
| privacy_policy | string | | No |
| prompt_public | boolean | | No |
-35
View File
@@ -1,35 +0,0 @@
"""Shared database fixtures for provider tests.
Provider tests live outside ``tests/unit_tests`` and therefore cannot use that
suite's SQLite fixtures. Keep this fixture scoped to ``providers`` so each
provider package can exercise real SQLAlchemy queries without a service
database or mocked sessions.
"""
from collections.abc import Iterator
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import Session
from models.base import TypeBase
@pytest.fixture
def sqlite3_session(request: pytest.FixtureRequest) -> Iterator[Session]:
"""Yield an isolated SQLite session with the parametrized model tables.
Pass the required ORM classes through indirect parametrization. The engine
is per-test so committed rows and identity-map state cannot leak between
provider packages.
"""
models: tuple[type[TypeBase], ...] = request.param
engine = create_engine("sqlite:///:memory:")
tables = [model.metadata.tables[model.__tablename__] for model in models]
TypeBase.metadata.create_all(engine, tables=tables)
try:
with Session(engine, expire_on_commit=False) as session:
yield session
finally:
engine.dispose()
-75
View File
@@ -1164,20 +1164,6 @@ class AgentComposerService:
agent.active_config_is_published = True
agent.updated_by = account_id
binding.current_snapshot_id = version.id
normal_draft = cls._get_agent_draft(
session=session,
tenant_id=tenant_id,
agent_id=agent.id,
draft_type=AgentConfigDraftType.DRAFT,
account_id=None,
)
if normal_draft is not None and cls._rebase_workflow_only_normal_draft(
agent=agent,
draft=normal_draft,
snapshot=version,
updated_by=account_id,
):
session.flush()
binding.updated_by = account_id
return binding
@@ -1762,47 +1748,6 @@ class AgentComposerService:
stmt = stmt.where(AgentConfigDraft.account_id.is_(None))
return session.scalar(stmt.order_by(AgentConfigDraft.updated_at.desc()).limit(1))
@classmethod
def get_or_create_normal_agent_draft(
cls,
*,
session: Session,
tenant_id: str,
agent: Agent,
created_by: str | None,
) -> AgentConfigDraft:
"""Resolve the shared Preview draft, rebasing inline agents when needed."""
return cls._get_or_create_agent_draft(
session=session,
tenant_id=tenant_id,
agent=agent,
draft_type=AgentConfigDraftType.DRAFT,
account_id=None,
created_by=created_by,
)
@staticmethod
def _rebase_workflow_only_normal_draft(
*,
agent: Agent,
draft: AgentConfigDraft,
snapshot: AgentConfigSnapshot,
updated_by: str | None,
) -> bool:
if (
agent.scope != AgentScope.WORKFLOW_ONLY
or draft.draft_type != AgentConfigDraftType.DRAFT
or draft.account_id is not None
or not agent.active_config_snapshot_id
or draft.base_snapshot_id == agent.active_config_snapshot_id
or snapshot.id != agent.active_config_snapshot_id
):
return False
draft.base_snapshot_id = snapshot.id
draft.config_snapshot = AgentSoulConfig.model_validate(snapshot.config_snapshot_dict)
draft.updated_by = updated_by
return True
@classmethod
def _get_or_create_agent_draft(
cls,
@@ -1822,26 +1767,6 @@ class AgentComposerService:
account_id=account_id,
)
if draft is not None:
if (
agent.scope == AgentScope.WORKFLOW_ONLY
and draft_type == AgentConfigDraftType.DRAFT
and draft.account_id is None
and agent.active_config_snapshot_id
and draft.base_snapshot_id != agent.active_config_snapshot_id
):
active_snapshot = cls._get_version_if_present(
session=session,
tenant_id=tenant_id,
agent_id=agent.id,
version_id=agent.active_config_snapshot_id,
)
if active_snapshot is not None and cls._rebase_workflow_only_normal_draft(
agent=agent,
draft=draft,
snapshot=active_snapshot,
updated_by=agent.updated_by or agent.created_by,
):
session.flush()
return draft
base_snapshot = cls._get_version_if_present(
session=session,
+23 -232
View File
@@ -9,13 +9,12 @@ import sqlalchemy as sa
from sqlalchemy import and_, func, or_, select
from sqlalchemy.orm import aliased
from configs import dify_config
from core.app.entities.app_invoke_entities import InvokeFrom
from libs.helper import convert_datetime_to_date, escape_like_pattern, to_timestamp
from models.agent import WorkflowAgentNodeBinding
from models.enums import CreatorUserRole, MessageStatus
from models.enums import MessageStatus
from models.model import App, Conversation, Message
from models.workflow import WorkflowNodeExecutionModel, WorkflowRun, WorkflowType
from models.workflow import WorkflowNodeExecutionModel, WorkflowRun
@dataclass(frozen=True)
@@ -581,26 +580,9 @@ class AgentObservabilityService:
def _load_daily_statistics(
self, *, app: App, agent_id: str, params: AgentStatisticsQueryParams, source_filter: AgentSourceFilter
) -> list[dict[str, Any]]:
rows: list[dict[str, Any]] = []
if source_filter.kind in {"all", "webapp"}:
rows.extend(self._load_webapp_daily_statistics(app=app, params=params, source_filter=source_filter))
if source_filter.kind in {"all", "workflow"}:
rows.extend(
self._load_workflow_daily_statistics(
app=app,
agent_id=agent_id,
params=params,
source_filter=source_filter,
)
)
return self._merge_daily_statistics(rows)
def _load_webapp_daily_statistics(
self, *, app: App, params: AgentStatisticsQueryParams, source_filter: AgentSourceFilter
) -> list[dict[str, Any]]:
converted_created_at = convert_datetime_to_date("m.created_at")
message_scope = self._statistics_webapp_message_scope_sql(source_filter)
message_scope = self._statistics_message_scope_sql(source_filter)
sql_query = f"""SELECT
{converted_created_at} AS date,
COUNT(m.id) AS message_count,
@@ -620,9 +602,20 @@ WHERE
args: dict[str, Any] = {
"tz": params.timezone,
"app_id": app.id,
"tenant_id": app.tenant_id,
"agent_id": agent_id,
"debugger": InvokeFrom.DEBUGGER,
}
if source_filter.invoke_from is not None:
args["source"] = source_filter.invoke_from
if source_filter.app_id:
args["source_app_id"] = source_filter.app_id
if source_filter.workflow_id:
args["workflow_id"] = source_filter.workflow_id
if source_filter.workflow_version:
args["workflow_version"] = source_filter.workflow_version
if source_filter.node_id:
args["node_id"] = source_filter.node_id
if params.start:
sql_query += " AND m.created_at >= :start"
args["start"] = params.start
@@ -634,160 +627,10 @@ WHERE
return [dict(row._mapping) for row in self._session.execute(sa.text(sql_query), args).all()]
@staticmethod
def _statistics_webapp_message_scope_sql(source_filter: AgentSourceFilter) -> str:
def _statistics_message_scope_sql(source_filter: AgentSourceFilter) -> str:
app_scope = "m.app_id = :app_id"
if source_filter.invoke_from is not None:
app_scope += " AND m.invoke_from = :source"
return app_scope
def _load_workflow_daily_statistics(
self,
*,
app: App,
agent_id: str,
params: AgentStatisticsQueryParams,
source_filter: AgentSourceFilter,
) -> list[dict[str, Any]]:
converted_run_created_at = convert_datetime_to_date("aru.created_at")
total_tokens = self._workflow_execution_metadata_numeric_sql(("total_tokens",), "BIGINT")
nested_total_tokens = self._workflow_execution_metadata_numeric_sql(
("agent_log", "agent_backend", "usage", "total_tokens"), "BIGINT"
)
total_price = self._workflow_execution_metadata_numeric_sql(("total_price",), "DECIMAL(65, 30)")
nested_total_price = self._workflow_execution_metadata_numeric_sql(
("agent_log", "agent_backend", "usage", "total_price"), "DECIMAL(65, 30)"
)
completion_tokens = self._workflow_execution_metadata_numeric_sql(
("agent_log", "agent_backend", "usage", "completion_tokens"), "BIGINT"
)
binding_filters = self._statistics_workflow_binding_filters_sql(source_filter)
run_date_filters = ""
args: dict[str, Any] = {
"tz": params.timezone,
"tenant_id": app.tenant_id,
"agent_id": agent_id,
"chat_workflow_type": WorkflowType.CHAT,
"end_user_role": CreatorUserRole.END_USER,
}
if source_filter.app_id:
args["source_app_id"] = source_filter.app_id
if source_filter.workflow_id:
args["workflow_id"] = source_filter.workflow_id
if source_filter.workflow_version:
args["workflow_version"] = source_filter.workflow_version
if source_filter.node_id:
args["node_id"] = source_filter.node_id
if params.start:
run_date_filters += " AND wr.created_at >= :start"
args["start"] = params.start
if params.end:
run_date_filters += " AND wr.created_at < :end"
args["end"] = params.end
run_query = f"""WITH agent_run_usage AS (
SELECT
wr.id,
wr.created_by_role,
wr.created_by,
wr.created_at,
COALESCE(SUM(COALESCE({total_tokens}, {nested_total_tokens}, 0)), 0) AS token_count,
COALESCE(SUM(COALESCE({total_price}, {nested_total_price}, 0)), 0) AS total_price,
COALESCE(SUM(COALESCE(wne.elapsed_time, 0)), 0) AS latency,
COALESCE(SUM(COALESCE({completion_tokens}, 0)), 0) AS answer_tokens
FROM workflow_runs wr
JOIN workflow_agent_node_bindings wanb
ON wanb.tenant_id = :tenant_id
AND wanb.agent_id = :agent_id
AND wanb.app_id = wr.app_id
AND wanb.workflow_id = wr.workflow_id
AND wanb.workflow_version = wr.version
{binding_filters}
JOIN workflow_node_executions wne
ON wne.workflow_run_id = wr.id
AND wne.node_id = wanb.node_id
WHERE wr.type != :chat_workflow_type{run_date_filters}
GROUP BY wr.id, wr.created_by_role, wr.created_by, wr.created_at
)
SELECT
{converted_run_created_at} AS date,
COUNT(aru.id) AS message_count,
COUNT(aru.id) AS conversation_count,
COUNT(DISTINCT CASE
WHEN aru.created_by_role = :end_user_role THEN aru.created_by
ELSE NULL
END) AS end_user_count,
COALESCE(SUM(aru.token_count), 0) AS token_count,
COALESCE(SUM(aru.total_price), 0) AS total_price,
COALESCE(AVG(aru.latency), 0) AS avg_latency,
COALESCE(SUM(aru.latency), 0) AS latency_sum,
COALESCE(SUM(aru.answer_tokens), 0) AS answer_tokens,
0 AS like_count
FROM agent_run_usage aru
GROUP BY date
ORDER BY date"""
rows = [dict(row._mapping) for row in self._session.execute(sa.text(run_query), args).all()]
rows.extend(
self._load_workflow_chat_daily_context(
app=app,
agent_id=agent_id,
params=params,
source_filter=source_filter,
)
)
return self._merge_daily_statistics(rows)
def _load_workflow_chat_daily_context(
self,
*,
app: App,
agent_id: str,
params: AgentStatisticsQueryParams,
source_filter: AgentSourceFilter,
) -> list[dict[str, Any]]:
converted_created_at = convert_datetime_to_date("m.created_at")
workflow_scope = self._statistics_workflow_message_scope_sql(source_filter)
sql_query = f"""SELECT
{converted_created_at} AS date,
COUNT(m.id) AS message_count,
COUNT(DISTINCT m.conversation_id) AS conversation_count,
COUNT(DISTINCT m.from_end_user_id) AS end_user_count,
COALESCE(SUM(COALESCE(m.message_tokens, 0) + COALESCE(m.answer_tokens, 0)), 0) AS token_count,
COALESCE(SUM(COALESCE(m.total_price, 0)), 0) AS total_price,
COALESCE(AVG(m.provider_response_latency), 0) AS avg_latency,
COALESCE(SUM(m.provider_response_latency), 0) AS latency_sum,
COALESCE(SUM(m.answer_tokens), 0) AS answer_tokens,
COUNT(mf.id) AS like_count
FROM messages m
LEFT JOIN message_feedbacks mf
ON mf.message_id = m.id AND mf.rating = 'like'
WHERE
{workflow_scope}"""
args: dict[str, Any] = {
"tz": params.timezone,
"tenant_id": app.tenant_id,
"agent_id": agent_id,
"chat_workflow_type": WorkflowType.CHAT,
}
if source_filter.app_id:
args["source_app_id"] = source_filter.app_id
if source_filter.workflow_id:
args["workflow_id"] = source_filter.workflow_id
if source_filter.workflow_version:
args["workflow_version"] = source_filter.workflow_version
if source_filter.node_id:
args["node_id"] = source_filter.node_id
if params.start:
sql_query += " AND m.created_at >= :start"
args["start"] = params.start
if params.end:
sql_query += " AND m.created_at < :end"
args["end"] = params.end
sql_query += " GROUP BY date ORDER BY date"
return [dict(row._mapping) for row in self._session.execute(sa.text(sql_query), args).all()]
@staticmethod
def _statistics_workflow_binding_filters_sql(source_filter: AgentSourceFilter) -> str:
workflow_binding_filters = []
if source_filter.app_id:
workflow_binding_filters.append("wanb.app_id = :source_app_id")
@@ -797,12 +640,8 @@ WHERE
workflow_binding_filters.append("wanb.workflow_version = :workflow_version")
if source_filter.node_id:
workflow_binding_filters.append("wanb.node_id = :node_id")
return f"AND {' AND '.join(workflow_binding_filters)}" if workflow_binding_filters else ""
@classmethod
def _statistics_workflow_message_scope_sql(cls, source_filter: AgentSourceFilter) -> str:
binding_filters = cls._statistics_workflow_binding_filters_sql(source_filter)
return f"""m.workflow_run_id IS NOT NULL
extra_workflow_filters = f"AND {' AND '.join(workflow_binding_filters)}" if workflow_binding_filters else ""
workflow_scope = f"""m.workflow_run_id IS NOT NULL
AND EXISTS (
SELECT 1
FROM workflow_runs wr
@@ -812,65 +651,17 @@ WHERE
AND wanb.app_id = wr.app_id
AND wanb.workflow_id = wr.workflow_id
AND wanb.workflow_version = wr.version
{binding_filters}
{extra_workflow_filters}
JOIN workflow_node_executions wne
ON wne.workflow_run_id = wr.id
AND wne.node_id = wanb.node_id
WHERE wr.id = m.workflow_run_id
AND wr.type = :chat_workflow_type
)"""
@staticmethod
def _workflow_execution_metadata_numeric_sql(path: tuple[str, ...], numeric_type: str) -> str:
if dify_config.DB_TYPE == "postgresql":
json_path = ",".join(path)
value = f"CAST(wne.execution_metadata AS JSONB) #>> '{{{json_path}}}'"
return f"CAST(NULLIF({value}, '') AS {numeric_type})"
if dify_config.DB_TYPE in {"mysql", "oceanbase", "seekdb"}:
json_path = "$." + ".".join(path)
mysql_numeric_type = "UNSIGNED" if numeric_type == "BIGINT" else numeric_type
value = f"JSON_UNQUOTE(JSON_EXTRACT(wne.execution_metadata, '{json_path}'))"
return f"CAST(NULLIF(NULLIF({value}, ''), 'null') AS {mysql_numeric_type})"
raise NotImplementedError(f"Unsupported database type: {dify_config.DB_TYPE}")
@staticmethod
def _merge_daily_statistics(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
merged: dict[Any, dict[str, Any]] = {}
weighted_latency: dict[Any, float] = {}
for row in rows:
date = row["date"]
target = merged.setdefault(
date,
{
"date": date,
"message_count": 0,
"conversation_count": 0,
"end_user_count": 0,
"token_count": 0,
"total_price": Decimal(0),
"avg_latency": 0.0,
"latency_sum": 0.0,
"answer_tokens": 0,
"like_count": 0,
},
)
message_count = int(row.get("message_count") or 0)
target["message_count"] += message_count
target["conversation_count"] += int(row.get("conversation_count") or 0)
target["end_user_count"] += int(row.get("end_user_count") or 0)
target["token_count"] += int(row.get("token_count") or 0)
target["total_price"] += Decimal(str(row.get("total_price") or 0))
target["latency_sum"] += float(row.get("latency_sum") or 0)
target["answer_tokens"] += int(row.get("answer_tokens") or 0)
target["like_count"] += int(row.get("like_count") or 0)
weighted_latency[date] = (
weighted_latency.get(date, 0.0) + float(row.get("avg_latency") or 0) * message_count
)
for date, row in merged.items():
message_count = int(row["message_count"])
row["avg_latency"] = weighted_latency[date] / message_count if message_count else 0.0
return sorted(merged.values(), key=lambda row: str(row["date"]))
if source_filter.kind == "webapp":
return app_scope
if source_filter.kind == "workflow":
return workflow_scope
return f"(({app_scope}) OR ({workflow_scope}))"
@staticmethod
def _build_charts(rows: list[dict[str, Any]]) -> dict[str, list[dict[str, Any]]]:
+23 -94
View File
@@ -430,11 +430,7 @@ class AgentRosterService:
agent.active_config_has_model = agent_soul_has_model(soul)
agent.active_config_is_published = False
self._session.flush()
self._get_or_create_agent_app_debug_conversation(
agent=agent,
account_id=account_id,
draft_type=AgentConfigDraftType.DEBUG_BUILD,
)
self._get_or_create_agent_app_debug_conversation(agent=agent, account_id=account_id)
return agent
def create_hidden_backing_app_for_workflow_agent(
@@ -531,11 +527,7 @@ class AgentRosterService:
self._session.flush()
return backing_app.id
def _get_or_create_agent_app_debug_conversation(
self, *, agent: Agent, account_id: str, draft_type: AgentConfigDraftType
) -> str:
"""Return the editor's conversation for one Agent draft surface."""
def _get_or_create_agent_app_debug_conversation(self, *, agent: Agent, account_id: str) -> str:
backing_app_id = self._ensure_workflow_agent_backing_app(agent=agent, account_id=account_id)
if not backing_app_id:
raise AgentNotFoundError()
@@ -545,7 +537,6 @@ class AgentRosterService:
AgentDebugConversation.tenant_id == agent.tenant_id,
AgentDebugConversation.agent_id == agent.id,
AgentDebugConversation.account_id == account_id,
AgentDebugConversation.draft_type == draft_type,
)
)
if mapping is not None:
@@ -579,7 +570,6 @@ class AgentRosterService:
agent_id=agent.id,
app_id=backing_app_id,
account_id=account_id,
draft_type=draft_type,
conversation_id=conversation_id,
)
)
@@ -587,15 +577,9 @@ class AgentRosterService:
return conversation_id
def get_or_create_agent_app_debug_conversation_id(
self,
*,
tenant_id: str,
agent_id: str,
account_id: str,
draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD,
commit: bool = True,
self, *, tenant_id: str, agent_id: str, account_id: str, commit: bool = True
) -> str:
"""Return the current editor's Build or Preview conversation for an Agent App."""
"""Return the current editor's debug conversation for an Agent App."""
agent = self._session.scalar(
select(Agent).where(
@@ -607,24 +591,13 @@ class AgentRosterService:
if agent is None:
raise AgentNotFoundError()
conversation_id = self._get_or_create_agent_app_debug_conversation(
agent=agent,
account_id=account_id,
draft_type=draft_type,
)
conversation_id = self._get_or_create_agent_app_debug_conversation(agent=agent, account_id=account_id)
if commit:
self._session.commit()
return conversation_id
def load_agent_app_debug_conversation_id(
self,
*,
tenant_id: str,
agent_id: str,
account_id: str,
draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD,
) -> str | None:
"""Return the editor's existing scoped conversation without creating or repairing rows."""
def load_agent_app_debug_conversation_id(self, *, tenant_id: str, agent_id: str, account_id: str) -> str | None:
"""Return the current editor's existing debug conversation without creating or repairing rows."""
return self._session.scalar(
select(Conversation.id)
@@ -633,7 +606,6 @@ class AgentRosterService:
AgentDebugConversation.tenant_id == tenant_id,
AgentDebugConversation.agent_id == agent_id,
AgentDebugConversation.account_id == account_id,
AgentDebugConversation.draft_type == draft_type,
AgentDebugConversation.app_id == Conversation.app_id,
Conversation.from_source == ConversationFromSource.CONSOLE,
Conversation.from_account_id == account_id,
@@ -654,21 +626,16 @@ class AgentRosterService:
)
def refresh_agent_app_debug_conversation_id(
self,
*,
tenant_id: str,
agent_id: str,
account_id: str,
draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD,
commit: bool = True,
self, *, tenant_id: str, agent_id: str, account_id: str, commit: bool = True
) -> str:
"""Start a new scoped console conversation for the current Agent App editor.
"""Start a new console debug conversation for the current Agent App editor.
If this account already has a mapping for the requested draft surface, the previous
If this account already has a debug conversation mapping, the previous
conversation is abandoned first: any ACTIVE conversation-owned Agent
runtime sessions for that old conversation are sent through best-effort
backend cleanup and then retired locally even when enqueueing fails. The
other draft surface is left untouched.
backend cleanup and then retired locally even when enqueueing fails.
The debug mapping is then repointed to the freshly created
conversation.
"""
agent = self._session.scalar(
@@ -696,7 +663,6 @@ class AgentRosterService:
AgentDebugConversation.tenant_id == tenant_id,
AgentDebugConversation.agent_id == agent_id,
AgentDebugConversation.account_id == account_id,
AgentDebugConversation.draft_type == draft_type,
)
)
if mapping is None:
@@ -706,7 +672,6 @@ class AgentRosterService:
agent_id=agent_id,
app_id=backing_app_id,
account_id=account_id,
draft_type=draft_type,
conversation_id=conversation_id,
)
)
@@ -718,7 +683,6 @@ class AgentRosterService:
tenant_id=tenant_id,
agent_id=agent_id,
account_id=account_id,
draft_type=draft_type,
app_id=previous_app_id or backing_app_id,
conversation_id=previous_conversation_id,
)
@@ -735,7 +699,6 @@ class AgentRosterService:
tenant_id: str,
agent_id: str,
account_id: str,
draft_type: AgentConfigDraftType,
app_id: str,
conversation_id: str,
) -> None:
@@ -764,8 +727,7 @@ class AgentRosterService:
session_snapshot=stored_session.session_snapshot,
runtime_layer_specs=stored_session.runtime_layer_specs,
idempotency_key=(
f"{tenant_id}:{agent_id}:{account_id}:{draft_type.value}:{conversation_id}:"
"debug-session-cleanup:"
f"{tenant_id}:{agent_id}:{account_id}:{conversation_id}:debug-session-cleanup:"
f"{stored_session.scope.agent_id}:"
f"{stored_session.scope.agent_config_snapshot_id or 'no-config'}:"
f"{stored_session.backend_run_id or 'no-run'}"
@@ -776,7 +738,6 @@ class AgentRosterService:
"conversation_id": stored_session.scope.conversation_id,
"agent_id": stored_session.scope.agent_id,
"agent_config_snapshot_id": stored_session.scope.agent_config_snapshot_id,
"draft_type": draft_type.value,
"previous_agent_backend_run_id": stored_session.backend_run_id,
},
)
@@ -811,14 +772,9 @@ class AgentRosterService:
)
def load_or_create_agent_app_debug_conversation_ids_by_agent_id(
self,
*,
tenant_id: str,
agents: list[Agent],
account_id: str,
draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD,
self, *, tenant_id: str, agents: list[Agent], account_id: str
) -> dict[str, str]:
"""Return per-account scoped conversations for a page of Agent Apps."""
"""Return per-account debug conversations for a page of Agent Apps."""
conversation_ids_by_agent_id: dict[str, str] = {}
changed = False
@@ -828,7 +784,6 @@ class AgentRosterService:
conversation_ids_by_agent_id[agent.id] = self._get_or_create_agent_app_debug_conversation(
agent=agent,
account_id=account_id,
draft_type=draft_type,
)
changed = True
if changed:
@@ -911,14 +866,16 @@ class AgentRosterService:
raise AgentNotFoundError()
return app
def _get_runtime_resolvable_agent(self, *, tenant_id: str, agent_id: str) -> Agent | None:
"""Load an Agent that is eligible to resolve to a runtime backing App.
def get_agent_runtime_app_model(self, *, tenant_id: str, agent_id: str) -> App:
"""Resolve the App that backs an Agent runtime surface.
Shared by the runtime resolver and the read-only authorization resolver
so both agree on what counts as a resolvable Agent.
Roster Agents use their public Agent App. Workflow-only Agents use a
hidden Agent App stored in ``backing_app_id`` so console chat/logs can
reuse the app runtime without exposing the resource in workspace app
lists.
"""
return self._session.scalar(
agent = self._session.scalar(
select(Agent)
.where(
Agent.tenant_id == tenant_id,
@@ -936,34 +893,6 @@ class AgentRosterService:
)
.limit(1)
)
def peek_authz_app_id(self, *, tenant_id: str, agent_id: str) -> str | None:
"""Resolve the App id whose access policy governs an Agent.
Roster Agents are governed by their own Agent App, while workflow-only
Agents are governed by their parent workflow App: the hidden runtime
backing App never receives a resource access policy, so it must not be
used for authorization. Stays read-only unlike
:meth:`get_agent_runtime_app_model`, this never materializes the hidden
backing App. Returns ``None`` when the Agent does not resolve, leaving
the caller to decide how to treat it.
"""
agent = self._get_runtime_resolvable_agent(tenant_id=tenant_id, agent_id=agent_id)
if agent is None:
return None
return agent.app_id
def get_agent_runtime_app_model(self, *, tenant_id: str, agent_id: str) -> App:
"""Resolve the App that backs an Agent runtime surface.
Roster Agents use their public Agent App. Workflow-only Agents use a
hidden Agent App stored in ``backing_app_id`` so console chat/logs can
reuse the app runtime without exposing the resource in workspace app
lists.
"""
agent = self._get_runtime_resolvable_agent(tenant_id=tenant_id, agent_id=agent_id)
if agent is None:
raise AgentNotFoundError()
should_commit_backing_app = agent.scope == AgentScope.WORKFLOW_ONLY and not agent.backing_app_id
-26
View File
@@ -19,7 +19,6 @@ from configs import dify_config
from constants.dsl_version import CURRENT_APP_DSL_VERSION
from core.file import remote_fetcher
from core.plugin.entities.plugin import PluginDependency
from core.rbac import RBACPermission
from core.trigger.constants import (
TRIGGER_PLUGIN_NODE_TYPE,
TRIGGER_SCHEDULE_NODE_TYPE,
@@ -44,9 +43,7 @@ from services.agent.dsl_service import AgentDslService, AgentPackage
from services.agent.workflow_publish_service import WorkflowAgentPublishService
from services.dsl_content import DSL_MAX_SIZE, dsl_content_size
from services.dsl_version import check_version_compatibility
from services.enterprise.rbac_service import RBACService
from services.entities.dsl_entities import CheckDependenciesResult, DslImportWarning, ImportMode, ImportStatus
from services.errors.account import NoPermissionError
from services.errors.app import WorkflowNotFoundError
from services.plugin.dependencies_analysis import DependenciesAnalysisService
from services.workflow_draft_variable_service import WorkflowDraftVariableService
@@ -304,9 +301,6 @@ class AppDslService:
error=f"Invalid YAML format: {str(e)}",
)
except NoPermissionError:
raise
except Exception as e:
logger.exception("Failed to import app")
return Import(
@@ -370,9 +364,6 @@ class AppDslService:
warnings=self._warnings,
)
except NoPermissionError:
raise
except Exception as e:
logger.exception("Error confirming import")
return Import(
@@ -404,21 +395,6 @@ class AppDslService:
leaked_dependencies=leaked_dependencies,
)
@staticmethod
def _ensure_agent_manage_permission(account: Account) -> None:
"""Importing an Agent DSL creates a roster Agent, which requires ``agent.manage``."""
if not dify_config.RBAC_ENABLED:
return
if account.current_tenant_id is None:
raise ValueError("Current tenant is not set")
allowed = RBACService.CheckAccess.check(
account.current_tenant_id,
account.id,
scene=RBACPermission.AGENT_MANAGE,
)
if not allowed:
raise NoPermissionError("Agent management permission is required to import an Agent App")
def _create_or_update_app(
self,
*,
@@ -439,8 +415,6 @@ class AppDslService:
if not app_mode:
raise ValueError("loss app mode")
app_mode = AppMode(app_mode)
if app_mode == AppMode.AGENT:
self._ensure_agent_manage_permission(account)
# Set icon type
icon_type_value = icon_type or app_data.get("icon_type")
+1 -7
View File
@@ -217,18 +217,12 @@ class BillingService:
return _billing_info_adapter.validate_python(billing_info)
@classmethod
def get_vector_space(cls, tenant_id: str, bypass_cache: bool = False) -> _VectorSpaceQuota:
def get_vector_space(cls, tenant_id: str) -> _VectorSpaceQuota:
params = {"tenant_id": tenant_id}
if bypass_cache:
params["bypass_cache"] = "true"
return _vector_space_quota_adapter.validate_python(
cls._send_request("GET", "/subscription/vector-space", params=params)
)
@classmethod
def invalidate_vector_space_cache(cls, tenant_id: str) -> None:
cls.get_vector_space(tenant_id, bypass_cache=True)
@classmethod
def get_tenant_feature_plan_usage_info(cls, tenant_id: str):
"""Deprecated: Use get_quota_info instead."""
+2 -10
View File
@@ -1974,10 +1974,7 @@ class DocumentService:
if data_source_info and "upload_file_id" in data_source_info:
file_id = data_source_info["upload_file_id"]
document_was_deleted.send(
document.id,
dataset_id=document.dataset_id,
doc_form=document.doc_form,
file_id=file_id,
document.id, dataset_id=document.dataset_id, doc_form=document.doc_form, file_id=file_id
)
session.delete(document)
@@ -2016,12 +2013,7 @@ class DocumentService:
# Dispatch cleanup task after commit to avoid lock contention
# Task cleans up segments, files, and vector indexes
if deleted_document_ids and doc_form is not None:
batch_clean_document_task.delay(
deleted_document_ids,
dataset_ref.dataset_id,
doc_form,
file_ids,
)
batch_clean_document_task.delay(deleted_document_ids, dataset_ref.dataset_id, doc_form, file_ids)
@staticmethod
def rename_document(dataset_id: str, document_id: str, name: str, session: Session) -> Document:
+8 -10
View File
@@ -5,7 +5,6 @@ from typing import Any
import httpx
from configs import dify_config
from core.helper.trace_id_helper import generate_traceparent_header
from services.errors.enterprise import (
EnterpriseAPIBadRequestError,
@@ -97,14 +96,12 @@ class BaseRequest:
logger.debug("Failed to generate traceparent header", exc_info=True)
with httpx.Client(mounts=mounts) as client:
# Callers that pass an explicit timeout keep it; everyone else gets the
# configured budget rather than httpx's implicit 5s default.
request_kwargs: dict[str, Any] = {
"json": json,
"params": params,
"headers": headers,
"timeout": timeout if timeout is not None else dify_config.ENTERPRISE_REQUEST_TIMEOUT,
}
# IMPORTANT:
# - In httpx, passing timeout=None disables timeouts (infinite) and overrides the library default.
# - To preserve httpx's default timeout behavior for existing call sites, only pass the kwarg when set.
request_kwargs: dict[str, Any] = {"json": json, "params": params, "headers": headers}
if timeout is not None:
request_kwargs["timeout"] = timeout
response = client.request(method, url, **request_kwargs)
@@ -209,8 +206,9 @@ class EnterpriseRequest(BaseRequest):
"json": json,
"params": params,
"headers": {"Content-Type": "application/json", cls.secret_key_header: cls.secret_key, **inner_headers},
"timeout": timeout if timeout is not None else dify_config.ENTERPRISE_RBAC_REQUEST_TIMEOUT,
}
if timeout is not None:
request_kwargs["timeout"] = timeout
response = client.request(method, url, **request_kwargs)
if not response.is_success:
cls._handle_error_response(response)
+13 -3
View File
@@ -317,6 +317,9 @@ _LEGACY_WORKSPACE_OWNER_KEYS: list[str] = [
"credential.use",
"credential.create",
"credential.manage",
"billing.view",
"billing.subscription.manage",
"billing.manage",
"app.acl.preview",
"app_library.access",
"app.create_and_management",
@@ -330,7 +333,6 @@ _LEGACY_WORKSPACE_OWNER_KEYS: list[str] = [
"snippets.management",
"tool.manage",
"mcp.manage",
"agent.manage",
]
_LEGACY_WORKSPACE_ADMIN_KEYS: list[str] = [
@@ -347,6 +349,9 @@ _LEGACY_WORKSPACE_ADMIN_KEYS: list[str] = [
"credential.use",
"credential.create",
"credential.manage",
"billing.view",
"billing.subscription.manage",
"billing.manage",
"app_library.access",
"app.create_and_management",
"app.tag.manage",
@@ -358,7 +363,6 @@ _LEGACY_WORKSPACE_ADMIN_KEYS: list[str] = [
"snippets.management",
"tool.manage",
"mcp.manage",
"agent.manage",
]
_LEGACY_WORKSPACE_EDITOR_KEYS: list[str] = [
@@ -373,7 +377,9 @@ _LEGACY_WORKSPACE_EDITOR_KEYS: list[str] = [
"dataset.external.connect",
"snippets.create_and_modify",
"tool.manage",
"agent.manage",
"billing.view",
"billing.subscription.manage",
"billing.manage",
]
_LEGACY_WORKSPACE_NORMAL_KEYS: list[str] = [
@@ -381,6 +387,9 @@ _LEGACY_WORKSPACE_NORMAL_KEYS: list[str] = [
"plugin.install",
"credential.use",
"app_library.access",
"billing.view",
"billing.subscription.manage",
"billing.manage",
]
_LEGACY_WORKSPACE_DATASET_OPERATOR_KEYS: list[str] = [
@@ -796,6 +805,7 @@ class RBACService:
data = _inner_call(
"GET",
f"{_INNER_PREFIX}/role-permissions/catalog",
params={"billing_enabled": dify_config.BILLING_ENABLED},
tenant_id=tenant_id,
account_id=account_id,
)
+1 -13
View File
@@ -93,20 +93,8 @@ class ExternalDatasetService:
raise ValueError(f"invalid endpoint: {endpoint} must start with http:// or https://")
else:
raise ValueError(f"invalid endpoint: {endpoint}")
# Send a minimal body shaped like the External Knowledge API retrieval contract so providers
# that require a JSON payload (e.g. RAGFlow) accept the validation probe instead of rejecting
# a body-less POST. Mirrors the request built in fetch_external_knowledge_retrieval.
validation_payload = {
"knowledge_id": "",
"query": "",
"retrieval_setting": {"top_k": 1, "score_threshold": 0.0},
}
try:
response = ssrf_proxy.post(
endpoint,
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
data=json.dumps(validation_payload),
)
response = ssrf_proxy.post(endpoint, headers={"Authorization": f"Bearer {api_key}"})
except Exception as e:
raise ValueError(f"failed to connect to the endpoint: {endpoint}") from e
if response.status_code == 502:
+1 -21
View File
@@ -75,12 +75,6 @@ class LicenseStatus(StrEnum):
LOST = "lost"
class DeploymentEdition(StrEnum):
COMMUNITY = "COMMUNITY"
ENTERPRISE = "ENTERPRISE"
CLOUD = "CLOUD"
class LicenseModel(FeatureResponseModel):
status: LicenseStatus = LicenseStatus.NONE
expired_at: str = ""
@@ -106,7 +100,6 @@ class WebAppAuthModel(FeatureResponseModel):
sso_config: WebAppAuthSSOModel = WebAppAuthSSOModel()
allow_email_code_login: bool = False
allow_email_password_login: bool = False
allow_public_access: bool = True
class KnowledgePipeline(FeatureResponseModel):
@@ -168,7 +161,6 @@ class PluginManagerModel(FeatureResponseModel):
class SystemFeatureModel(FeatureResponseModel):
deployment_edition: DeploymentEdition
enable_app_deploy: bool = False
sso_enforced_for_signin: bool = False
sso_enforced_for_signin_protocol: str = ""
@@ -193,7 +185,6 @@ class SystemFeatureModel(FeatureResponseModel):
enable_learn_app: bool = True
enable_step_by_step_tour: bool = False
rbac_enabled: bool = False
knowledge_fs_enabled: bool = False
class FeatureService:
@@ -259,7 +250,7 @@ class FeatureService:
@classmethod
def get_system_features(cls, is_authenticated: bool = False) -> SystemFeatureModel:
system_features = SystemFeatureModel(deployment_edition=cls._resolve_deployment_edition())
system_features = SystemFeatureModel()
system_features.rbac_enabled = dify_config.RBAC_ENABLED
cls._fulfill_system_params_from_env(system_features)
@@ -279,14 +270,6 @@ class FeatureService:
return system_features
@classmethod
def _resolve_deployment_edition(cls) -> DeploymentEdition:
if dify_config.EDITION == "CLOUD":
return DeploymentEdition.CLOUD
if dify_config.ENTERPRISE_ENABLED:
return DeploymentEdition.ENTERPRISE
return DeploymentEdition.COMMUNITY
@classmethod
def get_app_dsl_version(cls) -> str:
return CURRENT_APP_DSL_VERSION
@@ -300,13 +283,10 @@ class FeatureService:
system_features.is_allow_register = dify_config.ALLOW_REGISTER
system_features.is_allow_create_workspace = dify_config.ALLOW_CREATE_WORKSPACE
system_features.is_email_setup = dify_config.MAIL_TYPE is not None and dify_config.MAIL_TYPE != ""
system_features.enable_change_email = dify_config.ENABLE_CHANGE_EMAIL
system_features.enable_trial_app = dify_config.ENABLE_TRIAL_APP
system_features.enable_explore_banner = dify_config.ENABLE_EXPLORE_BANNER
system_features.enable_learn_app = dify_config.ENABLE_LEARN_APP
system_features.webapp_auth.allow_public_access = dify_config.WEBAPP_PUBLIC_ACCESS_ENABLED
system_features.enable_step_by_step_tour = dify_config.ENABLE_STEP_BY_STEP_TOUR
system_features.knowledge_fs_enabled = dify_config.KNOWLEDGE_FS_ENABLED
@classmethod
def _fulfill_trial_models_from_env(cls) -> list[str]:
-23
View File
@@ -141,29 +141,6 @@ class FileService:
blob = storage.load_once(upload_file_key)
return base64.b64encode(blob).decode()
def get_file_presigned_url(self, *, file_id: str, tenant_id: str) -> str:
"""Generate a direct storage URL for a tenant-owned upload file."""
with self._session_maker(expire_on_commit=False) as session:
upload_file = session.scalar(
select(UploadFile)
.where(
UploadFile.id == file_id,
UploadFile.tenant_id == tenant_id,
)
.limit(1)
)
if upload_file is None:
raise NotFound("File not found")
file_key = upload_file.key
content_type = upload_file.mime_type
return storage.generate_presigned_url(
file_key,
expires_in=dify_config.FILES_ACCESS_TIMEOUT,
content_type=content_type,
)
def upload_text(self, text: str, text_name: str, user_id: str, tenant_id: str) -> UploadFile:
if len(text_name) > 200:
text_name = text_name[:200]
-502
View File
@@ -1,502 +0,0 @@
"""Product-facing KnowledgeFS operation and authorization declarations.
This registry is Dify's explicit Console surface. Transport concerns live in
knowledge_fs_proxy so contract review does not require reading proxy mechanics.
"""
from __future__ import annotations
from typing import Final, Literal, NamedTuple
from core.rbac import RBACPermission
type KnowledgeFSMethod = Literal["DELETE", "GET", "PATCH", "POST", "PUT"]
type KnowledgeFSResponseKind = Literal["binary", "buffered", "stream"]
type KnowledgeFSRequiredScope = Literal["knowledge-spaces:read", "knowledge-spaces:write"]
type KnowledgeFSLegacyRole = Literal["reader", "dataset_editor", "admin"]
type KnowledgeFSErrorStatusMap = tuple[tuple[int, int], ...]
class KnowledgeFSOperation(NamedTuple):
operation_id: str
method: KnowledgeFSMethod
path: str
response_kind: KnowledgeFSResponseKind
required_scope: KnowledgeFSRequiredScope
rbac_permission: RBACPermission
legacy_role: KnowledgeFSLegacyRole
max_response_bytes: int
request_headers: tuple[str, ...]
response_headers: tuple[str, ...]
response_media_types: tuple[str, ...]
error_status_map: KnowledgeFSErrorStatusMap
def _console_operation(
operation_id: str,
method: KnowledgeFSMethod,
path: str,
*,
rbac_permission: RBACPermission,
legacy_role: KnowledgeFSLegacyRole,
max_response_bytes: int = 1_048_576,
request_headers: tuple[str, ...] = ("x-trace-id",),
response_kind: KnowledgeFSResponseKind = "buffered",
response_media_types: tuple[str, ...] = ("application/json",),
error_status_map: KnowledgeFSErrorStatusMap = ((401, 502), (403, 403)),
) -> KnowledgeFSOperation:
"""Declare one contract-pinned operation with an explicit Dify authorization policy."""
is_read = method == "GET"
return KnowledgeFSOperation(
operation_id=operation_id,
method=method,
path=path,
response_kind=response_kind,
required_scope="knowledge-spaces:read" if is_read else "knowledge-spaces:write",
rbac_permission=rbac_permission,
legacy_role=legacy_role,
max_response_bytes=max_response_bytes,
request_headers=request_headers,
response_headers=("x-trace-id",),
response_media_types=response_media_types,
error_status_map=error_status_map,
)
def _dataset_read_operation(operation_id: str, path: str) -> KnowledgeFSOperation:
"""Declare a dataset-readable buffered JSON operation."""
return _console_operation(
operation_id,
"GET",
path,
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
)
def _dataset_edit_operation(
operation_id: str,
method: KnowledgeFSMethod,
path: str,
*,
request_headers: tuple[str, ...] = ("x-trace-id",),
) -> KnowledgeFSOperation:
"""Declare a dataset-editable buffered JSON operation."""
return _console_operation(
operation_id,
method,
path,
rbac_permission=RBACPermission.DATASET_EDIT,
legacy_role="dataset_editor",
request_headers=request_headers,
)
def _external_source_operation(
operation_id: str,
method: KnowledgeFSMethod,
path: str,
*,
request_headers: tuple[str, ...] = ("x-trace-id",),
) -> KnowledgeFSOperation:
"""Declare a source-connection operation restricted to dataset editors."""
return _console_operation(
operation_id,
method,
path,
rbac_permission=RBACPermission.DATASET_EXTERNAL_CONNECT,
legacy_role="dataset_editor",
request_headers=request_headers,
)
KNOWLEDGE_FS_CONSOLE_OPERATIONS: Final[tuple[KnowledgeFSOperation, ...]] = (
_console_operation(
operation_id="listKnowledgeSpaces",
method="GET",
path="knowledge-spaces",
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
),
_console_operation(
operation_id="createKnowledgeSpace",
method="POST",
path="knowledge-spaces",
rbac_permission=RBACPermission.DATASET_CREATE_AND_MANAGEMENT,
legacy_role="dataset_editor",
),
_console_operation(
operation_id="getKnowledgeSpacesById",
method="GET",
path="knowledge-spaces/{id}",
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
),
_dataset_edit_operation("patchKnowledgeSpacesById", "PATCH", "knowledge-spaces/{id}"),
_dataset_edit_operation(
"deleteKnowledgeSpacesById",
"DELETE",
"knowledge-spaces/{id}",
request_headers=("idempotency-key", "x-trace-id"),
),
_dataset_read_operation("getKnowledgeSpacesByIdStats", "knowledge-spaces/{id}/stats"),
_console_operation(
operation_id="getKnowledgeSpacesByIdAccessPolicy",
method="GET",
path="knowledge-spaces/{id}/access-policy",
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
),
_console_operation(
operation_id="patchKnowledgeSpacesByIdAccessPolicy",
method="PATCH",
path="knowledge-spaces/{id}/access-policy",
rbac_permission=RBACPermission.DATASET_ACCESS_CONFIG,
legacy_role="admin",
),
_console_operation(
operation_id="getSourceProviders",
method="GET",
path="source-providers",
rbac_permission=RBACPermission.DATASET_EXTERNAL_CONNECT,
legacy_role="dataset_editor",
),
_console_operation(
operation_id="getKnowledgeSpacesByIdSourceConnections",
method="GET",
path="knowledge-spaces/{id}/source-connections",
rbac_permission=RBACPermission.DATASET_EXTERNAL_CONNECT,
legacy_role="dataset_editor",
),
_console_operation(
operation_id="postKnowledgeSpacesByIdSourceConnections",
method="POST",
path="knowledge-spaces/{id}/source-connections",
rbac_permission=RBACPermission.DATASET_EXTERNAL_CONNECT,
legacy_role="dataset_editor",
),
_external_source_operation(
"postKnowledgeSpacesByIdSourceConnectionsOauth",
"POST",
"knowledge-spaces/{id}/source-connections/oauth",
),
_external_source_operation("postSourceOauthCallback", "POST", "source-oauth/callback"),
_external_source_operation(
"getKnowledgeSpacesByIdSourceConnectionsByConnectionId",
"GET",
"knowledge-spaces/{id}/source-connections/{connectionId}",
),
_external_source_operation(
"deleteKnowledgeSpacesByIdSourceConnectionsByConnectionId",
"DELETE",
"knowledge-spaces/{id}/source-connections/{connectionId}",
),
_console_operation(
operation_id="postKnowledgeSpacesByIdSourceConnectionsByConnectionIdRefresh",
method="POST",
path="knowledge-spaces/{id}/source-connections/{connectionId}/refresh",
rbac_permission=RBACPermission.DATASET_EXTERNAL_CONNECT,
legacy_role="dataset_editor",
),
_console_operation(
operation_id="getKnowledgeSpacesByIdSources",
method="GET",
path="knowledge-spaces/{id}/sources",
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
),
_console_operation(
operation_id="postKnowledgeSpacesByIdSources",
method="POST",
path="knowledge-spaces/{id}/sources",
rbac_permission=RBACPermission.DATASET_EXTERNAL_CONNECT,
legacy_role="dataset_editor",
),
_external_source_operation(
"getKnowledgeSpacesByIdSourcesBySourceId",
"GET",
"knowledge-spaces/{id}/sources/{sourceId}",
),
_external_source_operation(
"patchKnowledgeSpacesByIdSourcesBySourceId",
"PATCH",
"knowledge-spaces/{id}/sources/{sourceId}",
),
_external_source_operation(
"deleteKnowledgeSpacesByIdSourcesBySourceId",
"DELETE",
"knowledge-spaces/{id}/sources/{sourceId}",
request_headers=("idempotency-key", "x-trace-id"),
),
_external_source_operation(
"putKnowledgeSpacesByIdSourcesBySourceIdCredentials",
"PUT",
"knowledge-spaces/{id}/sources/{sourceId}/credentials",
),
_external_source_operation(
"deleteKnowledgeSpacesByIdSourcesBySourceIdCredentials",
"DELETE",
"knowledge-spaces/{id}/sources/{sourceId}/credentials",
),
_external_source_operation(
"postKnowledgeSpacesByIdSourcesBySourceIdSync",
"POST",
"knowledge-spaces/{id}/sources/{sourceId}/sync",
request_headers=("idempotency-key", "x-trace-id"),
),
_console_operation(
operation_id="postKnowledgeSpacesByIdSourcesBySourceIdCrawlPreview",
method="POST",
path="knowledge-spaces/{id}/sources/{sourceId}/crawl-preview",
rbac_permission=RBACPermission.DATASET_EXTERNAL_CONNECT,
legacy_role="dataset_editor",
request_headers=("idempotency-key", "x-trace-id"),
),
_external_source_operation(
"postKnowledgeSpacesByIdSourcesBySourceIdWorkflowImports",
"POST",
"knowledge-spaces/{id}/sources/{sourceId}/workflow-imports",
request_headers=("idempotency-key", "x-trace-id"),
),
_external_source_operation(
"getKnowledgeSpacesByIdSourcesBySourceIdPages",
"GET",
"knowledge-spaces/{id}/sources/{sourceId}/pages",
),
_external_source_operation(
"getKnowledgeSpacesByIdSourcesBySourceIdFiles",
"GET",
"knowledge-spaces/{id}/sources/{sourceId}/files",
),
_external_source_operation(
"postKnowledgeSpacesByIdSourcesBySourceIdCrawl",
"POST",
"knowledge-spaces/{id}/sources/{sourceId}/crawl",
),
_external_source_operation(
"postKnowledgeSpacesByIdSourcesBySourceIdImport",
"POST",
"knowledge-spaces/{id}/sources/{sourceId}/import",
),
_external_source_operation(
"postKnowledgeSpacesByIdSourcesBySourceIdTest",
"POST",
"knowledge-spaces/{id}/sources/{sourceId}/test",
),
_external_source_operation(
"postKnowledgeSpacesByIdSourcesBySourceIdImportFiles",
"POST",
"knowledge-spaces/{id}/sources/{sourceId}/import-files",
),
_external_source_operation(
"postKnowledgeSpacesByIdSourcesBulk",
"POST",
"knowledge-spaces/{id}/sources/bulk",
request_headers=("idempotency-key", "x-trace-id"),
),
_external_source_operation(
"getKnowledgeSpacesByIdSourceWorkflows",
"GET",
"knowledge-spaces/{id}/source-workflows",
),
_console_operation(
operation_id="getKnowledgeSpacesByIdSourceWorkflowsByRunId",
method="GET",
path="knowledge-spaces/{id}/source-workflows/{runId}",
rbac_permission=RBACPermission.DATASET_EXTERNAL_CONNECT,
legacy_role="dataset_editor",
),
_external_source_operation(
"getKnowledgeSpacesByIdSourceWorkflowsByRunIdBulkItems",
"GET",
"knowledge-spaces/{id}/source-workflows/{runId}/bulk-items",
),
_console_operation(
operation_id="getKnowledgeSpacesByIdSourceWorkflowsByRunIdPages",
method="GET",
path="knowledge-spaces/{id}/source-workflows/{runId}/pages",
rbac_permission=RBACPermission.DATASET_EXTERNAL_CONNECT,
legacy_role="dataset_editor",
),
_console_operation(
operation_id="postKnowledgeSpacesByIdSourceWorkflowsByRunIdCancel",
method="POST",
path="knowledge-spaces/{id}/source-workflows/{runId}/cancel",
rbac_permission=RBACPermission.DATASET_EXTERNAL_CONNECT,
legacy_role="dataset_editor",
),
_console_operation(
operation_id="postKnowledgeSpacesByIdSourceWorkflowsByRunIdRetry",
method="POST",
path="knowledge-spaces/{id}/source-workflows/{runId}/retry",
rbac_permission=RBACPermission.DATASET_EXTERNAL_CONNECT,
legacy_role="dataset_editor",
),
_console_operation(
operation_id="postKnowledgeSpacesByIdSourceWorkflowsByRunIdSelection",
method="POST",
path="knowledge-spaces/{id}/source-workflows/{runId}/selection",
rbac_permission=RBACPermission.DATASET_EXTERNAL_CONNECT,
legacy_role="dataset_editor",
request_headers=("idempotency-key", "x-trace-id"),
),
_console_operation(
operation_id="getKnowledgeSpacesByIdSourcesBySourceIdSyncPolicy",
method="GET",
path="knowledge-spaces/{id}/sources/{sourceId}/sync-policy",
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
),
_console_operation(
operation_id="putKnowledgeSpacesByIdSourcesBySourceIdSyncPolicy",
method="PUT",
path="knowledge-spaces/{id}/sources/{sourceId}/sync-policy",
rbac_permission=RBACPermission.DATASET_EDIT,
legacy_role="dataset_editor",
),
_dataset_read_operation("getKnowledgeSpacesByIdDocuments", "knowledge-spaces/{id}/documents"),
_dataset_edit_operation("postKnowledgeSpacesByIdDocuments", "POST", "knowledge-spaces/{id}/documents"),
_dataset_edit_operation(
"deleteKnowledgeSpacesByIdDocumentsBulk",
"DELETE",
"knowledge-spaces/{id}/documents/bulk",
request_headers=("idempotency-key", "x-trace-id"),
),
_dataset_edit_operation(
"postKnowledgeSpacesByIdDocumentsBulk",
"POST",
"knowledge-spaces/{id}/documents/bulk",
),
_dataset_edit_operation(
"postKnowledgeSpacesByIdDocumentsBulkReindex",
"POST",
"knowledge-spaces/{id}/documents/bulk/reindex",
),
_dataset_read_operation(
"getKnowledgeSpacesByIdDocumentsByDocumentId",
"knowledge-spaces/{id}/documents/{documentId}",
),
_dataset_edit_operation(
"deleteKnowledgeSpacesByIdDocumentsByDocumentId",
"DELETE",
"knowledge-spaces/{id}/documents/{documentId}",
request_headers=("idempotency-key", "x-trace-id"),
),
_console_operation(
operation_id="getKnowledgeSpacesByIdLogicalDocuments",
method="GET",
path="knowledge-spaces/{id}/logical-documents",
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
),
_dataset_edit_operation(
"deleteKnowledgeSpacesByIdLogicalDocumentsByDocumentId",
"DELETE",
"knowledge-spaces/{id}/logical-documents/{documentId}",
request_headers=("idempotency-key", "x-trace-id"),
),
_dataset_read_operation(
"getKnowledgeSpacesByIdDocumentsByDocumentIdOutline",
"knowledge-spaces/{id}/documents/{documentId}/outline",
),
_console_operation(
operation_id="getKnowledgeSpacesByIdLogicalDocumentsByDocumentId",
method="GET",
path="knowledge-spaces/{id}/logical-documents/{documentId}",
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
),
_console_operation(
operation_id="getKnowledgeSpacesByIdDocumentsByDocumentIdRevisions",
method="GET",
path="knowledge-spaces/{id}/documents/{documentId}/revisions",
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
),
_dataset_edit_operation(
"postKnowledgeSpacesByIdDocumentsByDocumentIdRevisionsByRevisionRollback",
"POST",
"knowledge-spaces/{id}/documents/{documentId}/revisions/{revision}/rollback",
),
_dataset_edit_operation(
"patchKnowledgeSpacesByIdDocumentsByDocumentIdMetadata",
"PATCH",
"knowledge-spaces/{id}/documents/{documentId}/metadata",
),
_console_operation(
operation_id="getKnowledgeSpacesByIdDocumentsByDocumentIdRevisionsByRevisionChunks",
method="GET",
path="knowledge-spaces/{id}/documents/{documentId}/revisions/{revision}/chunks",
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
),
_dataset_read_operation(
"getKnowledgeSpacesByIdDocumentsByDocumentIdRevisionsByRevisionChunksByChunkId",
"knowledge-spaces/{id}/documents/{documentId}/revisions/{revision}/chunks/{chunkId}",
),
_dataset_edit_operation(
"postKnowledgeSpacesByIdDocumentsByDocumentIdRevisionsByRevisionChunksByChunkIdState",
"POST",
"knowledge-spaces/{id}/documents/{documentId}/revisions/{revision}/chunks/{chunkId}/state",
),
_console_operation(
operation_id="getKnowledgeSpacesByIdProcessingTasks",
method="GET",
path="knowledge-spaces/{id}/processing-tasks",
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
),
_dataset_read_operation(
"getKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasks",
"knowledge-spaces/{id}/documents/{documentId}/processing-tasks",
),
_dataset_read_operation(
"getKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasksByTaskId",
"knowledge-spaces/{id}/documents/{documentId}/processing-tasks/{taskId}",
),
_console_operation(
operation_id="getKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasksByTaskIdEvents",
method="GET",
path="knowledge-spaces/{id}/documents/{documentId}/processing-tasks/{taskId}/events",
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
max_response_bytes=67_108_864,
request_headers=("last-event-id", "x-trace-id"),
response_kind="stream",
response_media_types=("text/event-stream",),
),
_console_operation(
operation_id="deleteKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasksByTaskId",
method="DELETE",
path="knowledge-spaces/{id}/documents/{documentId}/processing-tasks/{taskId}",
rbac_permission=RBACPermission.DATASET_EDIT,
legacy_role="dataset_editor",
),
_console_operation(
operation_id="postKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasksByTaskIdRetry",
method="POST",
path="knowledge-spaces/{id}/documents/{documentId}/processing-tasks/{taskId}/retry",
rbac_permission=RBACPermission.DATASET_EDIT,
legacy_role="dataset_editor",
),
_dataset_read_operation(
"getKnowledgeSpacesByIdDocumentsByDocumentIdSettings",
"knowledge-spaces/{id}/documents/{documentId}/settings",
),
_dataset_edit_operation(
"putKnowledgeSpacesByIdDocumentsByDocumentIdSettings",
"PUT",
"knowledge-spaces/{id}/documents/{documentId}/settings",
),
_dataset_read_operation("getJobsById", "jobs/{id}"),
_dataset_edit_operation("deleteJobsById", "DELETE", "jobs/{id}"),
_dataset_edit_operation("postJobsByIdRetry", "POST", "jobs/{id}/retry"),
_dataset_read_operation("getDeletionJobsByJobId", "deletion-jobs/{jobId}"),
_dataset_edit_operation(
"postDeletionJobsByJobIdRetry",
"POST",
"deletion-jobs/{jobId}/retry",
request_headers=("idempotency-key", "x-trace-id"),
),
_dataset_read_operation("getBulkJobsById", "bulk-jobs/{id}"),
)
+68 -118
View File
@@ -1,32 +1,33 @@
"""Authorize and forward the explicitly enabled KnowledgeFS Console operations.
"""Transport-only forwarding for the explicitly enabled KnowledgeFS Console operations.
The dedicated request path uses Dify's shared SSRF policy, never follows redirects,
bounds buffered responses, and rejects compressed streaming responses.
KnowledgeFS owns the request and response contract. This module binds short-lived
account and workspace identities, enforces Dify's coarse workspace policy, and
normalizes transport failures. Dify deliberately maintains a small product-facing
operation registry instead of exposing the full upstream OpenAPI surface. The
dedicated request path uses Dify's shared SSRF policy, never follows redirects,
bounds buffered responses, and rejects compressed responses.
"""
from __future__ import annotations
from collections.abc import Iterable, Mapping
from dataclasses import dataclass
from datetime import UTC, datetime, timedelta
from http import HTTPStatus
from typing import NamedTuple, Protocol
from typing import Final, Literal, NamedTuple, Protocol
import httpx
import jwt
from configs import dify_config
from core.helper import ssrf_proxy
from core.rbac import RBACResourceScope
from core.rbac import RBACPermission, RBACResourceScope
from core.tools.errors import ToolSSRFError
from models import Account
from services.enterprise.rbac_service import RBACService
from services.knowledge_fs_operations import (
KNOWLEDGE_FS_CONSOLE_OPERATIONS,
KnowledgeFSMethod,
KnowledgeFSOperation,
KnowledgeFSResponseKind,
)
type KnowledgeFSMethod = Literal["DELETE", "GET", "PATCH", "POST", "PUT"]
type KnowledgeFSResponseKind = Literal["binary", "buffered", "stream"]
type KnowledgeFSRequiredScope = Literal["knowledge-spaces:read", "knowledge-spaces:write"]
_JWT_AUDIENCE = "knowledge-fs"
_JWT_ISSUER = "dify"
@@ -34,55 +35,56 @@ _JWT_TTL_SECONDS = 60
_MAX_BUFFERED_RESPONSE_BYTES = 1024 * 1024
class KnowledgeFSOperation(NamedTuple):
operation_id: str
method: KnowledgeFSMethod
path: str
response_kind: KnowledgeFSResponseKind
required_scope: KnowledgeFSRequiredScope
rbac_permission: RBACPermission
requires_dataset_editor: bool
max_response_bytes: int
request_headers: tuple[str, ...]
response_headers: tuple[str, ...]
response_media_types: tuple[str, ...]
KNOWLEDGE_FS_CONSOLE_OPERATIONS: Final[tuple[KnowledgeFSOperation, ...]] = (
KnowledgeFSOperation(
operation_id="listKnowledgeSpaces",
method="GET",
path="knowledge-spaces",
response_kind="buffered",
required_scope="knowledge-spaces:read",
rbac_permission=RBACPermission.DATASET_READONLY,
requires_dataset_editor=False,
max_response_bytes=1_048_576,
request_headers=("x-trace-id",),
response_headers=("x-trace-id",),
response_media_types=("application/json",),
),
KnowledgeFSOperation(
operation_id="createKnowledgeSpace",
method="POST",
path="knowledge-spaces",
response_kind="buffered",
required_scope="knowledge-spaces:write",
rbac_permission=RBACPermission.DATASET_CREATE_AND_MANAGEMENT,
requires_dataset_editor=True,
max_response_bytes=1_048_576,
request_headers=("x-trace-id",),
response_headers=("x-trace-id",),
response_media_types=("application/json",),
),
)
class KnowledgeFSUpstreamResponse(NamedTuple):
response: httpx.Response
response_kind: KnowledgeFSResponseKind
operation: KnowledgeFSOperation
_AUTHORIZATION_MARKER = object()
@dataclass(eq=False, frozen=True, init=False, slots=True)
class KnowledgeFSAuthorization:
"""Single-use forwarding capability created after Dify workspace policy checks.
Callers obtain this value from :func:`authorize_knowledge_fs_request`. Direct
construction and repeated forwarding are rejected before outbound I/O.
"""
account_id: str
tenant_id: str
operation: KnowledgeFSOperation
_used: bool
def __init__(
self,
account_id: str,
tenant_id: str,
operation: KnowledgeFSOperation,
*,
_marker: object | None = None,
) -> None:
if _marker is not _AUTHORIZATION_MARKER:
raise KnowledgeFSAccessDeniedError("KnowledgeFS authorization must be created by workspace authorization")
object.__setattr__(self, "account_id", account_id)
object.__setattr__(self, "tenant_id", tenant_id)
object.__setattr__(self, "operation", operation)
object.__setattr__(self, "_used", False)
def consume(self) -> tuple[str, str, KnowledgeFSOperation]:
"""Return the authorized principals and canonical operation exactly once.
Raises:
KnowledgeFSAccessDeniedError: The capability was already consumed.
"""
if self._used:
raise KnowledgeFSAccessDeniedError("KnowledgeFS authorization has already been used")
object.__setattr__(self, "_used", True)
return self.account_id, self.tenant_id, self.operation
class _RequestHeaders(Protocol):
def items(self) -> Iterable[tuple[str, str]]: ...
@@ -111,29 +113,20 @@ def authorize_knowledge_fs_request(
*,
account: Account,
tenant_id: str,
method: KnowledgeFSMethod,
path: str,
) -> KnowledgeFSAuthorization:
operation: KnowledgeFSOperation,
) -> None:
"""Enforce Dify's workspace policy before KFS performs resource authorization.
Args:
account: Authenticated Dify account with its current workspace role.
tenant_id: Current Dify workspace identifier.
method: Requested upstream HTTP method.
path: Requested relative KnowledgeFS path.
operation: Dify-maintained KnowledgeFS operation and policy metadata.
Raises:
KnowledgeFSRouteNotAllowedError: The method and path do not resolve to a declared operation.
KnowledgeFSAccessDeniedError: The account lacks a required legacy or enterprise permission.
Returns:
A request-scoped capability binding the authorized account, workspace, and operation.
"""
operation = get_knowledge_fs_operation(method, path)
if operation.legacy_role == "dataset_editor" and not account.is_dataset_editor:
raise KnowledgeFSAccessDeniedError("KnowledgeFS operation requires dataset edit access")
if operation.legacy_role == "admin" and not account.is_admin_or_owner:
raise KnowledgeFSAccessDeniedError("KnowledgeFS operation requires workspace administration access")
if operation.requires_dataset_editor and not account.is_dataset_editor:
raise KnowledgeFSAccessDeniedError("KnowledgeFS mutations require dataset edit access")
if not RBACService.CheckAccess.check(
tenant_id,
account.id,
@@ -141,7 +134,6 @@ def authorize_knowledge_fs_request(
resource_type=RBACResourceScope.DATASET.value,
):
raise KnowledgeFSAccessDeniedError("KnowledgeFS operation is denied by workspace RBAC")
return KnowledgeFSAuthorization(account.id, tenant_id, operation, _marker=_AUTHORIZATION_MARKER)
def proxy_knowledge_fs_request(
@@ -157,62 +149,20 @@ def proxy_knowledge_fs_request(
request_headers: _RequestHeaders | None = None,
) -> KnowledgeFSUpstreamResponse:
"""Authorize and forward one allowlisted KnowledgeFS request as a single use case."""
authorization = authorize_knowledge_fs_request(
operation = get_knowledge_fs_operation(method, path)
authorize_knowledge_fs_request(
account=account,
tenant_id=tenant_id,
method=method,
path=path,
operation=operation,
)
return proxy_authorized_knowledge_fs_request(
authorization=authorization,
accept=accept,
content_type=content_type,
query=query,
body=body,
request_headers=request_headers,
)
def proxy_authorized_knowledge_fs_request(
*,
authorization: KnowledgeFSAuthorization,
accept: str | None = None,
content_type: str | None = None,
query: bytes | None = None,
body: bytes | None = None,
request_headers: _RequestHeaders | None = None,
) -> KnowledgeFSUpstreamResponse:
"""Forward one request whose operation and workspace policy were already authorized.
This performs one outbound KnowledgeFS request and does not repeat Dify RBAC checks.
Args:
authorization: Request-scoped capability returned by :func:`authorize_knowledge_fs_request`.
accept: Original Accept header, when present.
content_type: Original request Content-Type header, when present.
query: Original encoded query string from the Console request.
body: Original request body, when present.
request_headers: Incoming headers; only names declared by the operation are forwarded.
Returns:
The bounded KnowledgeFS response together with its transport metadata.
Raises:
KnowledgeFSConfigurationError: The connection is incomplete or blocked by outbound policy.
KnowledgeFSRouteNotAllowedError: A forwarded request header is outside the operation contract.
KnowledgeFSTimeoutError: KnowledgeFS exceeds the configured timeout.
KnowledgeFSTransportError: The request fails or its response violates transport bounds.
"""
account_id, tenant_id, operation = authorization.consume()
incoming_request_headers = {name.lower(): value for name, value in (request_headers or {}).items()}
contract_request_headers = {
name: incoming_request_headers[name] for name in operation.request_headers if name in incoming_request_headers
}
return _forward_knowledge_fs_request(
account_id=account_id,
method=operation.method,
path=operation.path,
account_id=account.id,
method=method,
path=path,
tenant_id=tenant_id,
accept=accept,
content_type=content_type,
+11 -90
View File
@@ -13,7 +13,6 @@ from sqlalchemy import desc, select
from sqlalchemy.orm import Session, sessionmaker
from core.app.apps.message_generator import MessageGenerator
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity
from core.app.entities.task_entities import (
HumanInputRequiredResponse,
MessageReplaceStreamResponse,
@@ -85,6 +84,9 @@ def build_workflow_event_stream(
topic = MessageGenerator.get_response_topic(app_mode, workflow_run.id)
workflow_run_repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(session_maker)
node_execution_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository(session_maker)
message_context = (
_get_message_context(session_maker, workflow_run.id) if app_mode == AppMode.ADVANCED_CHAT else None
)
pause_entity: WorkflowPauseEntity | None = None
if workflow_run.status == WorkflowExecutionStatus.PAUSED:
@@ -95,38 +97,6 @@ def build_workflow_event_stream(
pause_entity = None
resumption_context = _load_resumption_context(pause_entity)
message_context: MessageContext | None = None
if app_mode == AppMode.ADVANCED_CHAT:
if workflow_run.status == WorkflowExecutionStatus.PAUSED:
if resumption_context is None:
raise AssertionError(
"WorkflowResumptionContext is required for advanced-chat snapshot replay, "
f"workflow_run_id={workflow_run.id}"
)
generate_entity = resumption_context.get_generate_entity()
if not isinstance(generate_entity, AdvancedChatAppGenerateEntity):
raise AssertionError(
"AdvancedChatAppGenerateEntity is required for advanced-chat snapshot replay, "
f"workflow_run_id={workflow_run.id}, generate_entity_type={type(generate_entity).__name__}"
)
if not generate_entity.conversation_id:
raise AssertionError(
f"conversation_id is required for advanced-chat snapshot replay, workflow_run_id={workflow_run.id}"
)
message_context = _get_message_context_by_conversation(
session_maker,
conversation_id=generate_entity.conversation_id,
workflow_run_id=workflow_run.id,
)
else:
# Compatibility fallback for non-suspended snapshot requests. This app-scoped lookup is not optimal;
# a dedicated index or stronger lookup key would be preferable.
message_context = _get_message_context_by_app(
session_maker,
app_id=app_id,
workflow_run_id=workflow_run.id,
)
node_snapshots = node_execution_repo.get_execution_snapshots_by_workflow_run(
tenant_id=tenant_id,
app_id=app_id,
@@ -205,68 +175,19 @@ def build_workflow_event_stream(
return _generate()
def _get_message_context_by_conversation(
session_maker: sessionmaker[Session],
*,
conversation_id: str,
workflow_run_id: str,
) -> MessageContext | None:
"""Look up a paused or suspended Advanced Chat snapshot message by conversation and workflow run.
Use this exact lookup after recovering ``conversation_id`` from persisted resumption context. Its predicates match
``message_workflow_run_id_idx``.
"""
def _get_message_context(session_maker: sessionmaker[Session], workflow_run_id: str) -> MessageContext | None:
with session_maker() as session:
stmt = (
select(Message)
.where(
Message.conversation_id == conversation_id,
Message.workflow_run_id == workflow_run_id,
)
.order_by(desc(Message.created_at))
.limit(1)
)
stmt = select(Message).where(Message.workflow_run_id == workflow_run_id).order_by(desc(Message.created_at))
message = session.scalar(stmt)
if message is None:
return None
return _to_message_context(message)
def _get_message_context_by_app(
session_maker: sessionmaker[Session],
*,
app_id: str,
workflow_run_id: str,
) -> MessageContext | None:
"""Look up a non-suspended or running Advanced Chat reconnect snapshot by app and workflow run.
This compatibility path applies only when no resumption context is expected. The app-scoped query is not optimal;
a dedicated index or stronger lookup key would be preferable.
"""
with session_maker() as session:
stmt = (
select(Message)
.where(
Message.app_id == app_id,
Message.workflow_run_id == workflow_run_id,
)
.order_by(desc(Message.created_at))
.limit(1)
created_at = int(message.created_at.timestamp()) if message.created_at else 0
return MessageContext(
conversation_id=message.conversation_id,
message_id=message.id,
created_at=created_at,
answer=message.answer,
)
message = session.scalar(stmt)
if message is None:
return None
return _to_message_context(message)
def _to_message_context(message: Message) -> MessageContext:
created_at = int(message.created_at.timestamp()) if message.created_at else 0
return MessageContext(
conversation_id=message.conversation_id,
message_id=message.id,
created_at=created_at,
answer=message.answer,
)
def _load_resumption_context(pause_entity: WorkflowPauseEntity | None) -> WorkflowResumptionContext | None:
+1 -12
View File
@@ -13,7 +13,6 @@ from core.tools.utils.web_reader_tool import get_image_upload_file_ids
from extensions.ext_storage import storage
from models.dataset import Dataset, DatasetMetadataBinding, DocumentSegment
from models.model import UploadFile
from tasks.refresh_billing_vector_space_task import schedule_billing_vector_space_refresh
logger = logging.getLogger(__name__)
@@ -22,12 +21,7 @@ BATCH_SIZE = 1000
@shared_task(queue="dataset")
def batch_clean_document_task(
document_ids: list[str],
dataset_id: str,
doc_form: str | None,
file_ids: list[str],
) -> None:
def batch_clean_document_task(document_ids: list[str], dataset_id: str, doc_form: str | None, file_ids: list[str]):
"""
Clean document when document deleted.
:param document_ids: document ids
@@ -46,7 +40,6 @@ def batch_clean_document_task(
index_node_ids: list[str] = []
segment_ids: list[str] = []
total_image_upload_file_ids: list[str] = []
dataset_tenant_id: str | None = None
try:
# ============ Step 1: Query segment and file data (short read-only transaction) ============
@@ -95,7 +88,6 @@ def batch_clean_document_task(
delete_summaries=True,
session=session,
)
dataset_tenant_id = dataset.tenant_id
except Exception:
logger.exception(
"Failed to clean vector index for dataset_id: %s, document_ids: %s, index_node_ids count: %d",
@@ -211,9 +203,6 @@ def batch_clean_document_task(
dataset_id,
)
if dataset_tenant_id is not None:
schedule_billing_vector_space_refresh(dataset_tenant_id)
end_at = time.perf_counter()
logger.info(
click.style(
-5
View File
@@ -24,7 +24,6 @@ from models.dataset import (
)
from models.model import UploadFile
from models.workflow import Workflow
from tasks.refresh_billing_vector_space_task import schedule_billing_vector_space_refresh
logger = logging.getLogger(__name__)
@@ -53,7 +52,6 @@ def clean_dataset_task(
"""
logger.info(click.style(f"Start clean dataset when dataset deleted: {dataset_id}", fg="green"))
start_at = time.perf_counter()
vector_cleanup_succeeded = False
with session_factory.create_session() as session:
try:
@@ -95,7 +93,6 @@ def clean_dataset_task(
try:
index_processor = IndexProcessorFactory(doc_form).init_index_processor()
index_processor.clean(dataset, None, with_keywords=True, delete_child_chunks=True, session=session)
vector_cleanup_succeeded = True
logger.info(click.style(f"Successfully cleaned vector database for dataset: {dataset_id}", fg="green"))
except Exception:
logger.exception(click.style(f"Failed to clean vector database for dataset {dataset_id}", fg="red"))
@@ -189,8 +186,6 @@ def clean_dataset_task(
session.execute(file_delete_stmt)
session.commit()
if vector_cleanup_succeeded:
schedule_billing_vector_space_refresh(dataset.tenant_id)
end_at = time.perf_counter()
logger.info(
click.style(
+1 -13
View File
@@ -11,18 +11,12 @@ from core.tools.utils.web_reader_tool import get_image_upload_file_ids
from extensions.ext_storage import storage
from models.dataset import Dataset, DatasetMetadataBinding, DocumentSegment, SegmentAttachmentBinding
from models.model import UploadFile
from tasks.refresh_billing_vector_space_task import schedule_billing_vector_space_refresh
logger = logging.getLogger(__name__)
@shared_task(queue="dataset")
def clean_document_task(
document_id: str,
dataset_id: str,
doc_form: str,
file_id: str | None,
) -> None:
def clean_document_task(document_id: str, dataset_id: str, doc_form: str, file_id: str | None):
"""
Clean document when document deleted.
:param document_id: document id
@@ -35,7 +29,6 @@ def clean_document_task(
logger.info(click.style(f"Start clean document when document deleted: {document_id}", fg="green"))
start_at = time.perf_counter()
total_attachment_files = []
vector_cleanup_succeeded = False
with session_factory.create_session() as session:
try:
@@ -44,7 +37,6 @@ def clean_document_task(
if not dataset:
raise Exception("Document has no dataset")
dataset_tenant_id = dataset.tenant_id
segments = session.scalars(select(DocumentSegment).where(DocumentSegment.document_id == document_id)).all()
# Use JOIN to fetch attachments with bindings in a single query
attachments_with_bindings = session.execute(
@@ -90,7 +82,6 @@ def clean_document_task(
delete_summaries=True,
session=session,
)
vector_cleanup_succeeded = True
except Exception:
logger.exception(
"Failed to clean vector / keyword index in clean_document_task, "
@@ -163,9 +154,6 @@ def clean_document_task(
)
)
if vector_cleanup_succeeded:
schedule_billing_vector_space_refresh(dataset_tenant_id)
end_at = time.perf_counter()
logger.info(
click.style(
@@ -1,59 +0,0 @@
import logging
from celery import shared_task
from opentelemetry import metrics
from configs import dify_config
from services.billing_service import BillingService
logger = logging.getLogger(__name__)
_MAX_RETRIES = 3
_RETRY_DELAY_SECONDS = 30
_refresh_counter = metrics.get_meter(__name__).create_counter(
"billing.vector_space_cache_refresh.count",
description="Number of billing vector-space cache refresh outcomes",
)
@shared_task(queue="dataset", bind=True, max_retries=_MAX_RETRIES, default_retry_delay=_RETRY_DELAY_SECONDS)
def refresh_billing_vector_space_task(self, tenant_id: str) -> None:
"""Refresh billing vector-space usage after vector cleanup has completed."""
if not dify_config.BILLING_ENABLED:
return
try:
BillingService.invalidate_vector_space_cache(tenant_id)
except Exception as exc:
if self.request.retries >= _MAX_RETRIES:
_refresh_counter.add(1, {"outcome": "exhausted"})
logger.exception(
"Billing vector-space cache refresh retry budget exhausted, tenant_id=%s",
tenant_id,
)
raise
_refresh_counter.add(1, {"outcome": "retry"})
logger.warning(
"Billing vector-space cache refresh failed, scheduling retry %d/%d, tenant_id=%s",
self.request.retries + 1,
_MAX_RETRIES,
tenant_id,
exc_info=True,
)
raise self.retry(exc=exc, countdown=_RETRY_DELAY_SECONDS * (2**self.request.retries))
_refresh_counter.add(1, {"outcome": "success"})
logger.info("Billing vector-space cache refreshed, tenant_id=%s", tenant_id)
def schedule_billing_vector_space_refresh(tenant_id: str) -> None:
"""Dispatch a best-effort billing refresh without changing cleanup status."""
if not dify_config.BILLING_ENABLED:
return
try:
refresh_billing_vector_space_task.delay(tenant_id)
except Exception:
_refresh_counter.add(1, {"outcome": "dispatch_failure"})
logger.exception("Failed to dispatch billing vector-space cache refresh, tenant_id=%s", tenant_id)
@@ -0,0 +1,179 @@
"""Testcontainers integration tests for controllers.console.app.app_import endpoints."""
from __future__ import annotations
from inspect import unwrap
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from flask import Flask
from controllers.console.app import app_import as app_import_module
from services.app_dsl_service import ImportStatus
class _Result:
def __init__(self, status: ImportStatus, app_id: str | None = "app-1"):
self.status = status
self.app_id = app_id
def model_dump(self, mode: str = "json"):
return {"status": self.status, "app_id": self.app_id}
def _install_features(monkeypatch: pytest.MonkeyPatch, enabled: bool) -> None:
features = SimpleNamespace(webapp_auth=SimpleNamespace(enabled=enabled))
monkeypatch.setattr(app_import_module.FeatureService, "get_system_features", lambda: features)
class TestAppImportApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask):
return flask_app_with_containers
def test_import_post_returns_failed_status(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportApi()
method = unwrap(api.post)
_install_features(monkeypatch, enabled=False)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.FAILED, app_id=None),
)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method(api, SimpleNamespace(id="u1"))
assert status == 400
assert response["status"] == ImportStatus.FAILED
def test_import_post_returns_pending_status(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportApi()
method = unwrap(api.post)
_install_features(monkeypatch, enabled=False)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.PENDING),
)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method(api, SimpleNamespace(id="u1"))
assert status == 202
assert response["status"] == ImportStatus.PENDING
def test_import_post_updates_webapp_auth_when_enabled(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportApi()
method = unwrap(api.post)
_install_features(monkeypatch, enabled=True)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.COMPLETED, app_id="app-123"),
)
update_access = MagicMock()
monkeypatch.setattr(app_import_module.EnterpriseService.WebAppAuth, "update_app_access_mode", update_access)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method(api, SimpleNamespace(id="u1"))
update_access.assert_called_once_with("app-123", "private")
assert status == 200
assert response["status"] == ImportStatus.COMPLETED
def test_import_post_commits_session_on_success(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportApi()
method = unwrap(api.post)
_install_features(monkeypatch, enabled=False)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.COMPLETED, app_id="app-123"),
)
fake_session = MagicMock()
fake_session.__enter__.return_value = fake_session
fake_session.__exit__.return_value = None
monkeypatch.setattr(app_import_module, "Session", lambda *_args, **_kwargs: fake_session)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method(api, SimpleNamespace(id="u1"))
fake_session.commit.assert_called_once_with()
fake_session.rollback.assert_not_called()
assert status == 200
assert response["status"] == ImportStatus.COMPLETED
def test_import_post_rolls_back_session_on_failure(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportApi()
method = unwrap(api.post)
_install_features(monkeypatch, enabled=False)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.FAILED, app_id=None),
)
fake_session = MagicMock()
fake_session.__enter__.return_value = fake_session
fake_session.__exit__.return_value = None
monkeypatch.setattr(app_import_module, "Session", lambda *_args, **_kwargs: fake_session)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method(api, SimpleNamespace(id="u1"))
fake_session.rollback.assert_called_once_with()
fake_session.commit.assert_not_called()
assert status == 400
assert response["status"] == ImportStatus.FAILED
class TestAppImportConfirmApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask):
return flask_app_with_containers
def test_import_confirm_returns_failed_status(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportConfirmApi()
method = unwrap(api.post)
monkeypatch.setattr(
app_import_module.AppDslService,
"confirm_import",
lambda *_args, **_kwargs: _Result(ImportStatus.FAILED),
)
with app.test_request_context("/console/api/apps/imports/import-1/confirm", method="POST"):
response, status = method(api, SimpleNamespace(id="u1"), import_id="import-1")
assert status == 400
assert response["status"] == ImportStatus.FAILED
class TestAppImportCheckDependenciesApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask):
return flask_app_with_containers
def test_import_check_dependencies_returns_result(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = app_import_module.AppImportCheckDependenciesApi()
method = unwrap(api.get)
monkeypatch.setattr(
app_import_module.AppDslService,
"check_dependencies",
lambda *_args, **_kwargs: SimpleNamespace(model_dump=lambda mode="json": {"leaked_dependencies": []}),
)
with app.test_request_context("/console/api/apps/imports/app-1/check-dependencies", method="GET"):
response, status = method(api, app_model=SimpleNamespace(id="app-1"))
assert status == 200
assert response["leaked_dependencies"] == []
@@ -1,4 +1,4 @@
"""Unit tests for OAuth controller endpoints."""
"""Testcontainers integration tests for OAuth controller endpoints."""
from __future__ import annotations
@@ -16,10 +16,15 @@ from controllers.console.auth.oauth import (
)
from libs.oauth import OAuthUserInfo, encode_oauth_state
from models.account import AccountStatus
from services.account_service import AccountService
from services.errors.account import AccountRegisterError
class TestGetOAuthProviders:
@pytest.fixture
def app(self, flask_app_with_containers: Flask):
return flask_app_with_containers
@pytest.mark.parametrize(
("github_config", "google_config", "expected_github", "expected_google"),
[
@@ -60,6 +65,10 @@ class TestOAuthLogin:
def resource(self):
return OAuthLogin()
@pytest.fixture
def app(self, flask_app_with_containers: Flask):
return flask_app_with_containers
@pytest.fixture
def mock_oauth_provider(self):
provider = MagicMock()
@@ -172,6 +181,10 @@ class TestOAuthCallback:
def resource(self):
return OAuthCallback()
@pytest.fixture
def app(self, flask_app_with_containers: Flask):
return flask_app_with_containers
@pytest.fixture
def oauth_setup(self):
"""Common OAuth setup for callback tests"""
@@ -250,12 +263,10 @@ class TestOAuthCallback:
@patch("controllers.console.auth.oauth.dify_config")
@patch("controllers.console.auth.oauth.get_oauth_providers")
@patch("controllers.console.auth.oauth.RegisterService")
@patch("controllers.console.auth.oauth.AccountService")
@patch("controllers.console.auth.oauth.redirect")
def test_invitation_comparison_is_case_insensitive(
self,
mock_redirect,
mock_account_service,
mock_register_service,
mock_get_providers,
mock_config,
@@ -269,20 +280,13 @@ class TestOAuthCallback:
)
mock_get_providers.return_value = {"github": oauth_setup["provider"]}
mock_register_service.is_valid_invite_token.return_value = True
mock_register_service.get_invitation_if_token_valid.return_value = {
"account": oauth_setup["account"],
"data": {"email": "user@example.com"},
"tenant": MagicMock(),
}
mock_account_service.login.return_value = oauth_setup["token_pair"]
mock_register_service.get_invitation_by_token.return_value = {"email": "user@example.com"}
state = encode_oauth_state(invite_token="invite123", timezone="Asia/Shanghai")
with app.test_request_context(f"/auth/oauth/github/callback?code=test_code&state={state}"):
resource.get("github")
mock_register_service.get_invitation_if_token_valid.assert_called_once_with(
None, None, "invite123", session=ANY
)
mock_register_service.get_invitation_by_token.assert_called_once_with(token="invite123")
mock_redirect.assert_called_once_with("http://localhost:3000/signin/invite-settings?invite_token=invite123")
@pytest.mark.parametrize(
@@ -444,6 +448,10 @@ class TestOAuthCallback:
class TestAccountGeneration:
@pytest.fixture
def app(self, flask_app_with_containers: Flask):
return flask_app_with_containers
@pytest.fixture
def user_info(self):
return OAuthUserInfo(id="123", name="Test User", email="test@example.com")
@@ -460,25 +468,39 @@ class TestAccountGeneration:
self,
mock_account_model,
mock_get_account,
app: Flask,
flask_req_ctx_with_containers,
user_info: OAuthUserInfo,
mock_account,
):
with app.test_request_context("/"):
# Test OpenID found
mock_account_model.get_by_openid.return_value = mock_account
result = _get_account_by_openid_or_email("github", user_info)
assert result == mock_account
mock_account_model.get_by_openid.assert_called_once_with("github", "123")
mock_get_account.assert_not_called()
# Test OpenID found
mock_account_model.get_by_openid.return_value = mock_account
result = _get_account_by_openid_or_email("github", user_info)
assert result == mock_account
mock_account_model.get_by_openid.assert_called_once_with("github", "123")
mock_get_account.assert_not_called()
# Test fallback to email lookup
mock_account_model.get_by_openid.return_value = None
mock_get_account.return_value = mock_account
# Test fallback to email lookup
mock_account_model.get_by_openid.return_value = None
mock_get_account.return_value = mock_account
result = _get_account_by_openid_or_email("github", user_info)
assert result == mock_account
mock_get_account.assert_called_once()
result = _get_account_by_openid_or_email("github", user_info)
assert result == mock_account
mock_get_account.assert_called_once()
def test_get_account_by_email_with_case_fallback_falls_back_to_lowercase(self):
"""Test that case fallback tries lowercase when exact match fails."""
mock_session = MagicMock()
first_result = MagicMock()
first_result.scalar_one_or_none.return_value = None
expected_account = MagicMock()
second_result = MagicMock()
second_result.scalar_one_or_none.return_value = expected_account
mock_session.execute.side_effect = [first_result, second_result]
result = AccountService.get_account_by_email_with_case_fallback("Case@Test.com", session=mock_session)
assert result is expected_account
assert mock_session.execute.call_count == 2
@pytest.mark.parametrize(
("allow_register", "existing_account", "should_create"),
@@ -1,14 +1,12 @@
"""Unit tests for password reset controller flows."""
"""Testcontainers integration tests for password reset authentication flows."""
from __future__ import annotations
from collections.abc import Generator
from contextlib import contextmanager
from unittest.mock import MagicMock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from sqlalchemy.orm import Session
from controllers.console.auth.error import (
EmailCodeError,
@@ -23,61 +21,47 @@ from controllers.console.auth.forgot_password import (
ForgotPasswordSendEmailApi,
)
from controllers.console.error import AccountNotFound, EmailSendIpLimitError
from models.account import Account, Tenant, TenantAccountJoin
from services.feature_service import DeploymentEdition, SystemFeatureModel
SQLITE_MODELS = (Account, Tenant, TenantAccountJoin)
@contextmanager
def _bind_database_session(session: Session) -> Generator[scoped_session[Session]]:
"""Bind the controller's session proxy to the SQLite test engine."""
database_session = scoped_session(sessionmaker(bind=session.get_bind(), expire_on_commit=False))
try:
with patch("controllers.console.auth.forgot_password.db.session", database_session):
yield database_session
finally:
database_session.remove()
@pytest.fixture(autouse=True)
def enable_password_login_wrappers(monkeypatch: pytest.MonkeyPatch) -> None:
"""Keep endpoint decorators deterministic without requiring the configured app database."""
monkeypatch.setattr("controllers.console.wraps.dify_config.EDITION", "CLOUD")
monkeypatch.setattr(
"controllers.console.wraps.FeatureService.get_system_features",
lambda: SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
),
)
from tests.test_containers_integration_tests.controllers.console.helpers import ensure_dify_setup
class TestForgotPasswordSendEmailApi:
"""Test cases for sending password reset emails."""
@pytest.mark.parametrize("sqlite_session", [SQLITE_MODELS], indirect=True)
@pytest.fixture
def app(self, flask_app_with_containers: Flask, db_session_with_containers: Session):
ensure_dify_setup(db_session_with_containers)
return flask_app_with_containers
@pytest.fixture
def mock_account(self):
"""Create mock account object."""
account = MagicMock()
account.email = "test@example.com"
account.name = "Test User"
return account
@patch("controllers.console.auth.forgot_password.AccountService.is_email_send_ip_limit")
@patch("controllers.console.auth.forgot_password.AccountService.get_account_by_email_with_case_fallback")
@patch("controllers.console.auth.forgot_password.AccountService.send_reset_password_email")
@patch("controllers.console.auth.forgot_password.FeatureService.get_system_features")
def test_send_reset_email_success(
self,
mock_get_features,
mock_send_email,
mock_get_account,
mock_is_ip_limit,
app: Flask,
sqlite_session: Session,
mock_account,
):
# Arrange
mock_is_ip_limit.return_value = False
mock_get_account.return_value = mock_account
mock_send_email.return_value = "reset_token_123"
mock_get_features.return_value.is_allow_register = True
# Act
with (
_bind_database_session(sqlite_session),
app.test_request_context(
"/forgot-password", method="POST", json={"email": "test@example.com", "language": "en-US"}
),
with app.test_request_context(
"/forgot-password", method="POST", json={"email": "test@example.com", "language": "en-US"}
):
api = ForgotPasswordSendEmailApi()
response = api.post()
@@ -114,17 +98,20 @@ class TestForgotPasswordSendEmailApi:
(None, "en-US"), # Defaults to en-US when not provided
],
)
@pytest.mark.parametrize("sqlite_session", [SQLITE_MODELS], indirect=True)
@patch("controllers.console.auth.forgot_password.AccountService.is_email_send_ip_limit")
@patch("controllers.console.auth.forgot_password.AccountService.get_account_by_email_with_case_fallback")
@patch("controllers.console.auth.forgot_password.AccountService.send_reset_password_email")
@patch("controllers.console.auth.forgot_password.FeatureService.get_system_features")
def test_send_reset_email_language_handling(
self,
mock_get_features,
mock_send_email,
mock_get_account,
mock_is_ip_limit,
app: Flask,
mock_account,
language_input,
expected_language,
app: Flask,
sqlite_session: Session,
):
"""
Test password reset email with different language preferences.
@@ -135,14 +122,13 @@ class TestForgotPasswordSendEmailApi:
"""
# Arrange
mock_is_ip_limit.return_value = False
mock_get_account.return_value = mock_account
mock_send_email.return_value = "token"
mock_get_features.return_value.is_allow_register = True
# Act
with (
_bind_database_session(sqlite_session),
app.test_request_context(
"/forgot-password", method="POST", json={"email": "test@example.com", "language": language_input}
),
with app.test_request_context(
"/forgot-password", method="POST", json={"email": "test@example.com", "language": language_input}
):
api = ForgotPasswordSendEmailApi()
api.post()
@@ -155,6 +141,11 @@ class TestForgotPasswordSendEmailApi:
class TestForgotPasswordCheckApi:
"""Test cases for verifying password reset codes."""
@pytest.fixture
def app(self, flask_app_with_containers: Flask, db_session_with_containers: Session):
ensure_dify_setup(db_session_with_containers)
return flask_app_with_containers
@patch("controllers.console.auth.forgot_password.AccountService.is_forgot_password_error_rate_limit")
@patch("controllers.console.auth.forgot_password.AccountService.get_reset_password_data")
@patch("controllers.console.auth.forgot_password.AccountService.revoke_reset_password_token")
@@ -162,10 +153,10 @@ class TestForgotPasswordCheckApi:
@patch("controllers.console.auth.forgot_password.AccountService.reset_forgot_password_error_rate_limit")
def test_verify_code_success(
self,
mock_reset_rate_limit: MagicMock,
mock_generate_token: MagicMock,
mock_revoke_token: MagicMock,
mock_get_data: MagicMock,
mock_reset_rate_limit,
mock_generate_token,
mock_revoke_token,
mock_get_data,
mock_is_rate_limit,
app: Flask,
):
@@ -209,10 +200,10 @@ class TestForgotPasswordCheckApi:
@patch("controllers.console.auth.forgot_password.AccountService.reset_forgot_password_error_rate_limit")
def test_verify_code_preserves_token_email_case(
self,
mock_reset_rate_limit: MagicMock,
mock_generate_token: MagicMock,
mock_revoke_token: MagicMock,
mock_get_data: MagicMock,
mock_reset_rate_limit,
mock_generate_token,
mock_revoke_token,
mock_get_data,
mock_is_rate_limit,
app: Flask,
):
@@ -334,15 +325,33 @@ class TestForgotPasswordCheckApi:
class TestForgotPasswordResetApi:
"""Test cases for resetting password with verified token."""
@pytest.mark.parametrize("sqlite_session", [SQLITE_MODELS], indirect=True)
@pytest.fixture
def app(self, flask_app_with_containers: Flask, db_session_with_containers: Session):
ensure_dify_setup(db_session_with_containers)
return flask_app_with_containers
@pytest.fixture
def mock_account(self):
"""Create mock account object."""
account = MagicMock()
account.email = "test@example.com"
account.name = "Test User"
return account
@patch("controllers.console.auth.forgot_password.AccountService.get_reset_password_data")
@patch("controllers.console.auth.forgot_password.AccountService.revoke_reset_password_token")
@patch("controllers.console.auth.forgot_password.AccountService.get_account_by_email_with_case_fallback")
@patch("controllers.console.auth.forgot_password.db")
@patch("controllers.console.auth.forgot_password.TenantService.get_join_tenants")
def test_reset_password_success(
self,
mock_revoke_token: MagicMock,
mock_get_data: MagicMock,
mock_get_tenants,
mock_db,
mock_get_account,
mock_revoke_token,
mock_get_data,
app: Flask,
sqlite_session: Session,
mock_account,
):
"""
Test successful password reset.
@@ -354,39 +363,25 @@ class TestForgotPasswordResetApi:
"""
# Arrange
mock_get_data.return_value = {"email": "test@example.com", "phase": "reset"}
mock_get_account.return_value = mock_account
mock_db.session.merge.return_value = mock_account
mock_get_tenants.return_value = [MagicMock()]
# Act
with _bind_database_session(sqlite_session) as database_session:
account = Account(name="Test User", email="test@example.com")
tenant = Tenant(name="Test Workspace")
database_session.add_all([account, tenant])
database_session.flush()
database_session.add(TenantAccountJoin(tenant_id=tenant.id, account_id=account.id))
database_session.commit()
account_id = account.id
with app.test_request_context(
"/forgot-password/resets",
method="POST",
json={
"token": "valid_token",
"new_password": "NewPass123!",
"password_confirm": "NewPass123!",
},
):
api = ForgotPasswordResetApi()
response = api.post()
updated_account = database_session.get(Account, account_id)
with app.test_request_context(
"/forgot-password/resets",
method="POST",
json={"token": "valid_token", "new_password": "NewPass123!", "password_confirm": "NewPass123!"},
):
api = ForgotPasswordResetApi()
response = api.post()
# Assert
assert response["result"] == "success"
mock_revoke_token.assert_called_once_with("valid_token")
assert updated_account is not None
assert updated_account.password is not None
assert updated_account.password_salt is not None
def test_reset_password_mismatch(self, app: Flask):
@patch("controllers.console.auth.forgot_password.AccountService.get_reset_password_data")
def test_reset_password_mismatch(self, mock_get_data, app: Flask):
"""
Test password reset with mismatched passwords.
@@ -394,6 +389,9 @@ class TestForgotPasswordResetApi:
- PasswordMismatchError is raised when passwords don't match
- No password update occurs
"""
# Arrange
mock_get_data.return_value = {"email": "test@example.com", "phase": "reset"}
# Act & Assert
with app.test_request_context(
"/forgot-password/resets",
@@ -447,12 +445,10 @@ class TestForgotPasswordResetApi:
with pytest.raises(InvalidTokenError):
api.post()
@pytest.mark.parametrize("sqlite_session", [SQLITE_MODELS], indirect=True)
@patch("controllers.console.auth.forgot_password.AccountService.get_reset_password_data")
@patch("controllers.console.auth.forgot_password.AccountService.revoke_reset_password_token")
def test_reset_password_account_not_found(
self, mock_revoke_token, mock_get_data, app: Flask, sqlite_session: Session
):
@patch("controllers.console.auth.forgot_password.AccountService.get_account_by_email_with_case_fallback")
def test_reset_password_account_not_found(self, mock_get_account, mock_revoke_token, mock_get_data, app: Flask):
"""
Test password reset for non-existent account.
@@ -461,15 +457,13 @@ class TestForgotPasswordResetApi:
"""
# Arrange
mock_get_data.return_value = {"email": "nonexistent@example.com", "phase": "reset"}
mock_get_account.return_value = None
# Act & Assert
with (
_bind_database_session(sqlite_session),
app.test_request_context(
"/forgot-password/resets",
method="POST",
json={"token": "token", "new_password": "NewPass123!", "password_confirm": "NewPass123!"},
),
with app.test_request_context(
"/forgot-password/resets",
method="POST",
json={"token": "token", "new_password": "NewPass123!", "password_confirm": "NewPass123!"},
):
api = ForgotPasswordResetApi()
with pytest.raises(AccountNotFound):
@@ -9,9 +9,7 @@ from flask import Flask
from sqlalchemy.orm import Session
from werkzeug.exceptions import Forbidden
from configs import dify_config
from controllers.web.site import AppSiteApi, WebAppSiteResponse, WebModelConfigResponse
from extensions.storage.storage_type import StorageType
from models import Tenant, TenantStatus
from models.account import TenantCustomConfigDict
from models.model import App, AppMode, AppModelConfig, CustomizeTokenStrategy, EndUser, Site
@@ -98,39 +96,6 @@ class TestAppSiteApi:
assert result["plan"] == "basic"
assert result["enable_site"] is True
@patch("controllers.web.site.FileService.get_file_presigned_url")
@patch("controllers.web.site.FeatureService.get_features")
def test_image_icon_uses_s3_presigned_url(
self,
mock_features: MagicMock,
mock_get_file_presigned_url: MagicMock,
app: Flask,
db_session_with_containers: Session,
) -> None:
app.config["RESTX_MASK_HEADER"] = "X-Fields"
tenant = _create_tenant(db_session_with_containers)
app_model = _create_app(db_session_with_containers, tenant.id)
site = _create_site(db_session_with_containers, app_model.id)
site.icon_type = "image"
site.icon = "11111111-1111-4111-8111-111111111111"
db_session_with_containers.commit()
end_user = _end_user(tenant.id, app_model.id)
mock_features.return_value = FeatureModel(can_replace_logo=False)
mock_get_file_presigned_url.return_value = "https://s3.example.com/icon.png?signature=test"
with (
patch.object(dify_config, "EDITION", "CLOUD"),
patch.object(dify_config, "STORAGE_TYPE", StorageType.S3),
app.test_request_context("/site"),
):
result = AppSiteApi().get(app_model, end_user)
assert result["site"]["icon_url"] == "https://s3.example.com/icon.png?signature=test"
mock_get_file_presigned_url.assert_called_once_with(
file_id="11111111-1111-4111-8111-111111111111",
tenant_id=tenant.id,
)
def test_missing_site_raises_forbidden(self, app: Flask, db_session_with_containers: Session) -> None:
app.config["RESTX_MASK_HEADER"] = "X-Fields"
tenant = _create_tenant(db_session_with_containers)
@@ -1,38 +1,22 @@
"""Unit tests for API token caching and SQLite-backed token lookup."""
from __future__ import annotations
from collections.abc import Iterator
from datetime import datetime
from unittest.mock import MagicMock, patch
from uuid import uuid4
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from werkzeug.exceptions import Unauthorized
import services.api_token_service as api_token_service_module
from models.engine import db
from models.model import ApiToken
from services.api_token_service import ApiTokenCache, CachedApiToken
@pytest.fixture
def api_token_db() -> Iterator[Session]:
"""Provide the production database extension with an isolated SQLite token table."""
app = Flask(__name__)
app.config["SQLALCHEMY_DATABASE_URI"] = "sqlite:///:memory:"
db.init_app(app)
with app.app_context():
ApiToken.__table__.create(db.engine)
with Session(db.engine, expire_on_commit=False) as session:
yield session
class TestQueryTokenFromDb:
def test_should_return_api_token_and_cache_when_token_exists(self, api_token_db: Session) -> None:
def test_should_return_api_token_and_cache_when_token_exists(
self, flask_app_with_containers: Flask, db_session_with_containers
):
tenant_id = str(uuid4())
app_id = str(uuid4())
token_value = f"app-test-{uuid4()}"
@@ -43,8 +27,8 @@ class TestQueryTokenFromDb:
api_token.tenant_id = tenant_id
api_token.type = "app"
api_token.token = token_value
api_token_db.add(api_token)
api_token_db.commit()
db_session_with_containers.add(api_token)
db_session_with_containers.commit()
with (
patch.object(api_token_service_module.ApiTokenCache, "set") as mock_cache_set,
@@ -57,7 +41,9 @@ class TestQueryTokenFromDb:
mock_cache_set.assert_called_once()
mock_record_usage.assert_called_once_with(token_value, "app")
def test_should_cache_null_and_raise_unauthorized_when_token_not_found(self, api_token_db: Session) -> None:
def test_should_cache_null_and_raise_unauthorized_when_token_not_found(
self, flask_app_with_containers: Flask, db_session_with_containers
):
with (
patch.object(api_token_service_module.ApiTokenCache, "set") as mock_cache_set,
patch.object(api_token_service_module, "record_token_usage") as mock_record_usage,
@@ -414,7 +414,6 @@ class TestFeatureService:
mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True
mock_config.ENABLE_SOCIAL_OAUTH_LOGIN = False
mock_config.ENABLE_COLLABORATION_MODE = False
mock_config.ENABLE_CHANGE_EMAIL = True
mock_config.ALLOW_REGISTER = True
mock_config.ALLOW_CREATE_WORKSPACE = True
mock_config.MAIL_TYPE = "smtp"
@@ -617,7 +616,6 @@ class TestFeatureService:
mock_config.ENABLE_EMAIL_CODE_LOGIN = False
mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True
mock_config.ENABLE_SOCIAL_OAUTH_LOGIN = True
mock_config.ENABLE_CHANGE_EMAIL = True
mock_config.ALLOW_REGISTER = False
mock_config.ALLOW_CREATE_WORKSPACE = False
mock_config.MAIL_TYPE = None
@@ -1,7 +1,6 @@
"""Unit tests for human-input test delivery with SQLite-backed member lookup."""
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from uuid import uuid4
@@ -20,9 +19,7 @@ from core.workflow.human_input_adapter import (
)
from graphon.runtime import VariablePool
from models.account import Account, TenantAccountJoin
from models.engine import db
from services import human_input_delivery_test_service as service_module
from services.feature_service import FeatureModel
from services.human_input_delivery_test_service import (
DeliveryTestContext,
DeliveryTestEmailRecipient,
@@ -93,14 +90,8 @@ class TestDeliveryTestRegistry:
with pytest.raises(DeliveryTestUnsupportedError, match="Delivery method does not support test send."):
registry.dispatch(context=context, method=method)
def test_default(self) -> None:
app = Flask(__name__)
app.config["SQLALCHEMY_DATABASE_URI"] = "sqlite:///:memory:"
db.init_app(app)
with app.app_context():
registry = DeliveryTestRegistry.default()
def test_default(self, flask_app_with_containers: Flask, db_session_with_containers: Session):
registry = DeliveryTestRegistry.default()
assert len(registry._handlers) == 1
assert isinstance(registry._handlers[0], EmailDeliveryTestHandler)
@@ -121,24 +112,24 @@ class TestEmailDeliveryTestHandler:
handler = EmailDeliveryTestHandler(session_factory=engine)
assert handler._session_factory.kw["bind"] == engine
def test_supports(self, sqlite_engine: Engine) -> None:
handler = EmailDeliveryTestHandler(session_factory=sqlite_engine)
def test_supports(self):
handler = EmailDeliveryTestHandler(session_factory=MagicMock())
method = EmailDeliveryMethod(config=_make_valid_email_config())
assert handler.supports(method) is True
assert handler.supports(MagicMock()) is False
def test_send_test_unsupported_method(self, sqlite_engine: Engine) -> None:
handler = EmailDeliveryTestHandler(session_factory=sqlite_engine)
def test_send_test_unsupported_method(self):
handler = EmailDeliveryTestHandler(session_factory=MagicMock())
with pytest.raises(DeliveryTestUnsupportedError):
handler.send_test(context=MagicMock(), method=MagicMock())
def test_send_test_feature_disabled(self, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None:
def test_send_test_feature_disabled(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
service_module.FeatureService,
"get_features",
lambda _tenant_id, **_kwargs: FeatureModel(human_input_email_delivery_enabled=False),
lambda _tenant_id, **_kwargs: SimpleNamespace(human_input_email_delivery_enabled=False),
)
handler = EmailDeliveryTestHandler(session_factory=sqlite_engine)
handler = EmailDeliveryTestHandler(session_factory=MagicMock())
context = DeliveryTestContext(
tenant_id="t1", app_id="a1", node_id="n1", node_title="title", rendered_content="content"
)
@@ -147,15 +138,15 @@ class TestEmailDeliveryTestHandler:
with pytest.raises(DeliveryTestError, match="Email delivery is not available"):
handler.send_test(context=context, method=method)
def test_send_test_mail_not_inited(self, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None:
def test_send_test_mail_not_inited(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
service_module.FeatureService,
"get_features",
lambda _id, **_kwargs: FeatureModel(human_input_email_delivery_enabled=True),
lambda _id, **_kwargs: SimpleNamespace(human_input_email_delivery_enabled=True),
)
monkeypatch.setattr(service_module.mail, "is_inited", lambda: False)
handler = EmailDeliveryTestHandler(session_factory=sqlite_engine)
handler = EmailDeliveryTestHandler(session_factory=MagicMock())
context = DeliveryTestContext(
tenant_id="t1", app_id="a1", node_id="n1", node_title="title", rendered_content="content"
)
@@ -164,15 +155,15 @@ class TestEmailDeliveryTestHandler:
with pytest.raises(DeliveryTestError, match="Mail client is not initialized."):
handler.send_test(context=context, method=method)
def test_send_test_no_recipients(self, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None:
def test_send_test_no_recipients(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
service_module.FeatureService,
"get_features",
lambda _id, **_kwargs: FeatureModel(human_input_email_delivery_enabled=True),
lambda _id, **_kwargs: SimpleNamespace(human_input_email_delivery_enabled=True),
)
monkeypatch.setattr(service_module.mail, "is_inited", lambda: True)
handler = EmailDeliveryTestHandler(session_factory=sqlite_engine)
handler = EmailDeliveryTestHandler(session_factory=MagicMock())
handler._resolve_recipients = MagicMock(return_value=[])
context = DeliveryTestContext(
@@ -183,18 +174,18 @@ class TestEmailDeliveryTestHandler:
with pytest.raises(DeliveryTestError, match="No recipients configured"):
handler.send_test(context=context, method=method)
def test_send_test_success(self, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None:
def test_send_test_success(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
service_module.FeatureService,
"get_features",
lambda _id, **_kwargs: FeatureModel(human_input_email_delivery_enabled=True),
lambda _id, **_kwargs: SimpleNamespace(human_input_email_delivery_enabled=True),
)
monkeypatch.setattr(service_module.mail, "is_inited", lambda: True)
mock_mail_send = MagicMock()
monkeypatch.setattr(service_module.mail, "send", mock_mail_send)
monkeypatch.setattr(service_module, "render_email_template", lambda t, s: f"RENDERED_{t}")
handler = EmailDeliveryTestHandler(session_factory=sqlite_engine)
handler = EmailDeliveryTestHandler(session_factory=MagicMock())
handler._resolve_recipients = MagicMock(return_value=["test@example.com"])
variable_pool = VariablePool()
@@ -219,11 +210,11 @@ class TestEmailDeliveryTestHandler:
assert kwargs["to"] == "test@example.com"
assert "RENDERED_Subj" in kwargs["subject"]
def test_send_test_sanitizes_subject(self, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None:
def test_send_test_sanitizes_subject(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
service_module.FeatureService,
"get_features",
lambda _id, **_kwargs: FeatureModel(human_input_email_delivery_enabled=True),
lambda _id, **_kwargs: SimpleNamespace(human_input_email_delivery_enabled=True),
)
monkeypatch.setattr(service_module.mail, "is_inited", lambda: True)
mock_mail_send = MagicMock()
@@ -234,7 +225,7 @@ class TestEmailDeliveryTestHandler:
lambda template, substitutions: template.replace("{{ recipient_email }}", substitutions["recipient_email"]),
)
handler = EmailDeliveryTestHandler(session_factory=sqlite_engine)
handler = EmailDeliveryTestHandler(session_factory=MagicMock())
handler._resolve_recipients = MagicMock(return_value=["test@example.com"])
context = DeliveryTestContext(
@@ -258,8 +249,8 @@ class TestEmailDeliveryTestHandler:
_, kwargs = mock_mail_send.call_args
assert kwargs["subject"] == "Notice BCC:test@example.com"
def test_resolve_recipients_external(self, sqlite_engine: Engine) -> None:
handler = EmailDeliveryTestHandler(session_factory=sqlite_engine)
def test_resolve_recipients_external(self):
handler = EmailDeliveryTestHandler(session_factory=MagicMock())
method = EmailDeliveryMethod(
config=EmailDeliveryConfig(
recipients=EmailRecipients(
@@ -271,18 +262,19 @@ class TestEmailDeliveryTestHandler:
)
assert handler._resolve_recipients(tenant_id="t1", method=method) == ["ext@example.com"]
@pytest.mark.parametrize("sqlite_session", [(Account, TenantAccountJoin)], indirect=True)
def test_resolve_recipients_member(self, sqlite_engine: Engine, sqlite_session: Session) -> None:
def test_resolve_recipients_member(self, flask_app_with_containers: Flask, db_session_with_containers: Session):
tenant_id = str(uuid4())
account = Account(name="Test User", email="member@example.com")
sqlite_session.add(account)
sqlite_session.commit()
db_session_with_containers.add(account)
db_session_with_containers.commit()
join = TenantAccountJoin(tenant_id=tenant_id, account_id=account.id)
sqlite_session.add(join)
sqlite_session.commit()
db_session_with_containers.add(join)
db_session_with_containers.commit()
handler = EmailDeliveryTestHandler(session_factory=sqlite_engine)
from extensions.ext_database import db
handler = EmailDeliveryTestHandler(session_factory=db.engine)
method = EmailDeliveryMethod(
config=EmailDeliveryConfig(
recipients=EmailRecipients(items=[MemberRecipient(reference_id=account.id)], include_bound_group=False),
@@ -292,20 +284,23 @@ class TestEmailDeliveryTestHandler:
)
assert handler._resolve_recipients(tenant_id=tenant_id, method=method) == ["member@example.com"]
@pytest.mark.parametrize("sqlite_session", [(Account, TenantAccountJoin)], indirect=True)
def test_resolve_recipients_whole_workspace(self, sqlite_engine: Engine, sqlite_session: Session) -> None:
def test_resolve_recipients_whole_workspace(
self, flask_app_with_containers: Flask, db_session_with_containers: Session
):
tenant_id = str(uuid4())
account1 = Account(name="User 1", email=f"u1-{uuid4()}@example.com")
account2 = Account(name="User 2", email=f"u2-{uuid4()}@example.com")
sqlite_session.add_all([account1, account2])
sqlite_session.commit()
db_session_with_containers.add_all([account1, account2])
db_session_with_containers.commit()
for acc in [account1, account2]:
join = TenantAccountJoin(tenant_id=tenant_id, account_id=acc.id)
sqlite_session.add(join)
sqlite_session.commit()
db_session_with_containers.add(join)
db_session_with_containers.commit()
handler = EmailDeliveryTestHandler(session_factory=sqlite_engine)
from extensions.ext_database import db
handler = EmailDeliveryTestHandler(session_factory=db.engine)
method = EmailDeliveryMethod(
config=EmailDeliveryConfig(
recipients=EmailRecipients(items=[], include_bound_group=True),
@@ -316,8 +311,8 @@ class TestEmailDeliveryTestHandler:
recipients = handler._resolve_recipients(tenant_id=tenant_id, method=method)
assert set(recipients) == {account1.email, account2.email}
def test_query_workspace_member_emails_empty_ids(self, sqlite_engine: Engine) -> None:
handler = EmailDeliveryTestHandler(session_factory=sqlite_engine)
def test_query_workspace_member_emails_empty_ids(self):
handler = EmailDeliveryTestHandler(session_factory=MagicMock())
assert handler._query_workspace_member_emails(tenant_id="t1", user_ids=[]) == {}
def test_build_substitutions(self):
@@ -212,17 +212,6 @@ def test_generate_specs_include_console_contract_shapes_for_schema_migration(tmp
assert {"type": "null"} in app_detail_nullable_schema["anyOf"]
assert schemas["RecommendedAppInfoResponse"]["properties"]["icon_url"]["readOnly"] is True
assert schemas["InstalledAppInfoResponse"]["properties"]["icon_url"]["readOnly"] is True
assert _response_schema(paths["/apps/{app_id}"]["get"])["$ref"] == "#/components/schemas/AppDetailWithSite"
app_model_config = schemas["AppDetailWithSite"]["properties"]["model_config"]
assert {"$ref": "#/components/schemas/AppModelConfigResponse"} in app_model_config["anyOf"]
app_detail = schemas["AppDetail"]
assert "mode" in app_detail["properties"]
assert "mode_compatible_with_agent" not in app_detail["properties"]
sync_draft_workflow = schemas["SyncDraftWorkflowResponse"]
assert _response_schema(paths["/apps/{app_id}/workflows/draft"]["post"])["$ref"] == (
"#/components/schemas/SyncDraftWorkflowResponse"
)
assert sync_draft_workflow["properties"]["updated_at"]["type"] == "integer"
tool_icon_schema = schemas["ExploreAppMetaResponse"]["properties"]["tool_icons"]["additionalProperties"]
assert {"type": "string"} in tool_icon_schema["anyOf"]
assert {"additionalProperties": True, "type": "object"} in tool_icon_schema["anyOf"]
@@ -1,24 +1,15 @@
"""SQLite-backed tests for the reset-encrypt-key-pair CLI command (#35396).
"""Unit tests for the reset-encrypt-key-pair CLI command (#35396).
The command must purge every table that stores ciphertext encrypted with the
tenant's asymmetric key, otherwise stale rows cause downstream API failures
such as `/console/api/workspaces/current/tool-providers` returning 500.
Tests bind the command-owned transaction to the fixture engine and assert the
committed state rather than inspecting fabricated ``Session.execute`` calls.
"""
from types import SimpleNamespace
import pytest
from sqlalchemy import select
from sqlalchemy.orm import Session
from unittest.mock import MagicMock, patch
import commands
from commands import system as system_commands
from core.tools.entities.tool_entities import ApiProviderSchemaType
from graphon.model_runtime.entities.model_entities import ModelType
from models import Tenant
from models.provider import Provider, ProviderModel, ProviderType
from models.provider import Provider, ProviderModel
from models.tools import ApiToolProvider, BuiltinToolProvider, MCPToolProvider
@@ -30,60 +21,17 @@ def _invoke_reset() -> int:
return 0
TENANT_ID = "11111111-1111-1111-1111-111111111111"
OTHER_TENANT_ID = "11111111-1111-1111-1111-111111111112"
USER_ID = "22222222-2222-2222-2222-222222222222"
def _tenant(tenant_id: str, *, name: str = "Test tenant") -> Tenant:
tenant = Tenant(name=name, encrypt_public_key="old-key")
tenant.id = tenant_id
return tenant
def _encrypted_rows(tenant_id: str, *, suffix: str = "1") -> tuple[object, ...]:
"""Build one persisted credential-bearing row for every purge target."""
return (
Provider(tenant_id=tenant_id, provider_name=f"provider-{suffix}"),
ProviderModel(
tenant_id=tenant_id,
provider_name=f"provider-{suffix}",
model_name=f"model-{suffix}",
model_type=ModelType.LLM,
),
BuiltinToolProvider(
name=f"builtin-credential-{suffix}",
tenant_id=tenant_id,
user_id=USER_ID,
provider=f"builtin-{suffix}",
encrypted_credentials="ciphertext",
),
ApiToolProvider(
name=f"api-{suffix}",
icon="icon",
schema="{}",
schema_type_str=ApiProviderSchemaType.OPENAPI,
user_id=USER_ID,
tenant_id=tenant_id,
description="description",
tools_str="[]",
credentials_str="{}",
),
MCPToolProvider(
name=f"mcp-{suffix}",
server_identifier=f"server-{suffix}",
server_url="ciphertext",
server_url_hash=f"hash-{suffix}",
icon=None,
tenant_id=tenant_id,
user_id=USER_ID,
encrypted_credentials="ciphertext",
),
)
def _bind_command_to_sqlite(monkeypatch: pytest.MonkeyPatch, session: Session) -> None:
monkeypatch.setattr(system_commands, "db", SimpleNamespace(engine=session.get_bind()))
def _delete_targets(session_mock: MagicMock) -> list:
"""Extract the model class targeted by each `delete(...)` call on the session."""
targets = []
for call in session_mock.execute.call_args_list:
stmt = call.args[0]
# `delete(Foo)` constructs a `Delete` statement whose entity is `Foo`.
try:
targets.append(stmt.table.name)
except AttributeError:
targets.append(repr(stmt))
return targets
def test_reset_aborts_when_not_self_hosted(monkeypatch, capsys):
@@ -96,73 +44,65 @@ def test_reset_aborts_when_not_self_hosted(monkeypatch, capsys):
assert "only for SELF_HOSTED" in captured.out
@pytest.mark.parametrize(
"sqlite_session",
[(Tenant, Provider, ProviderModel, BuiltinToolProvider, ApiToolProvider, MCPToolProvider)],
indirect=True,
)
def test_reset_purges_provider_and_tool_tables_for_each_tenant(
monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str], sqlite_session: Session
) -> None:
def test_reset_purges_provider_and_tool_tables_for_each_tenant(monkeypatch, capsys):
"""The command must purge LLM provider rows AND every tool provider table
that stores ciphertext encrypted under the tenant key (#35396)."""
monkeypatch.setattr(system_commands.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setattr(system_commands, "generate_key_pair", lambda tenant_id: f"new-key-{tenant_id}")
_bind_command_to_sqlite(monkeypatch, sqlite_session)
tenant = _tenant(TENANT_ID)
other_tenant = _tenant(OTHER_TENANT_ID, name="Other tenant")
system_provider = Provider(
tenant_id=TENANT_ID,
provider_name="system-provider",
provider_type=ProviderType.SYSTEM,
)
sqlite_session.add_all((tenant, other_tenant, system_provider, *_encrypted_rows(TENANT_ID)))
sqlite_session.commit()
fake_tenant = MagicMock(id="tenant-abc", encrypt_public_key="old-key")
session = MagicMock()
session.scalars.return_value.all.return_value = [fake_tenant]
exit_code = _invoke_reset()
fake_sessionmaker = MagicMock()
fake_sessionmaker.begin.return_value.__enter__.return_value = session
fake_sessionmaker.begin.return_value.__exit__.return_value = False
with (
patch.object(system_commands, "db", MagicMock()),
patch.object(system_commands, "sessionmaker", return_value=fake_sessionmaker),
):
exit_code = _invoke_reset()
captured = capsys.readouterr()
assert exit_code == 0
assert TENANT_ID in captured.out
assert "tenant-abc" in captured.out
sqlite_session.expire_all()
assert sqlite_session.get(Tenant, TENANT_ID).encrypt_public_key == f"new-key-{TENANT_ID}"
assert sqlite_session.get(Tenant, OTHER_TENANT_ID).encrypt_public_key == f"new-key-{OTHER_TENANT_ID}"
assert sqlite_session.scalars(select(Provider).where(Provider.provider_type == ProviderType.CUSTOM)).all() == []
assert sqlite_session.scalars(select(ProviderModel)).all() == []
assert sqlite_session.scalars(select(BuiltinToolProvider)).all() == []
assert sqlite_session.scalars(select(ApiToolProvider)).all() == []
assert sqlite_session.scalars(select(MCPToolProvider)).all() == []
assert (
sqlite_session.scalar(select(Provider).where(Provider.provider_type == ProviderType.SYSTEM)) is system_provider
)
# New key pair generated and assigned.
assert fake_tenant.encrypt_public_key == "new-key-tenant-abc"
# Every encrypted-credential table should have been purged for this tenant.
table_names = _delete_targets(session)
expected = {
Provider.__tablename__,
ProviderModel.__tablename__,
BuiltinToolProvider.__tablename__,
ApiToolProvider.__tablename__,
MCPToolProvider.__tablename__,
}
assert expected.issubset(set(table_names)), f"missing purges: expected {expected}, got {table_names}"
@pytest.mark.parametrize(
"sqlite_session",
[(Tenant, Provider, ProviderModel, BuiltinToolProvider, ApiToolProvider, MCPToolProvider)],
indirect=True,
)
def test_reset_iterates_all_tenants(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
def test_reset_iterates_all_tenants(monkeypatch, capsys):
"""Multi-tenant deployments must purge every tenant, not just the first."""
monkeypatch.setattr(system_commands.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setattr(system_commands, "generate_key_pair", lambda tenant_id: f"new-key-{tenant_id}")
_bind_command_to_sqlite(monkeypatch, sqlite_session)
tenant_ids = [f"11111111-1111-1111-1111-{index:012d}" for index in range(3)]
tenants = [_tenant(tenant_id, name=f"Tenant {index}") for index, tenant_id in enumerate(tenant_ids)]
for index, tenant in enumerate(tenants):
sqlite_session.add(tenant)
sqlite_session.add_all(_encrypted_rows(tenant.id, suffix=str(index)))
sqlite_session.commit()
tenants = [MagicMock(id=f"tenant-{i}", encrypt_public_key="old") for i in range(3)]
session = MagicMock()
session.scalars.return_value.all.return_value = tenants
assert _invoke_reset() == 0
fake_sessionmaker = MagicMock()
fake_sessionmaker.begin.return_value.__enter__.return_value = session
fake_sessionmaker.begin.return_value.__exit__.return_value = False
sqlite_session.expire_all()
persisted_tenants = sqlite_session.scalars(select(Tenant).order_by(Tenant.id)).all()
assert [tenant.encrypt_public_key for tenant in persisted_tenants] == [
f"new-key-{tenant_id}" for tenant_id in tenant_ids
]
for model in (Provider, ProviderModel, BuiltinToolProvider, ApiToolProvider, MCPToolProvider):
assert sqlite_session.scalars(select(model)).all() == []
with (
patch.object(system_commands, "db", MagicMock()),
patch.object(system_commands, "sessionmaker", return_value=fake_sessionmaker),
):
_invoke_reset()
# Five purges per tenant × 3 tenants = 15 execute calls.
assert session.execute.call_count == 15
for tenant in tenants:
assert tenant.encrypt_public_key == f"new-key-{tenant.id}"
+23 -37
View File
@@ -37,7 +37,6 @@ os.environ.setdefault("STORAGE_TYPE", "opendal")
from core.db.session_factory import configure_session_factory, session_factory
from extensions import ext_redis
from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole
from models.base import TypeBase
@@ -149,45 +148,32 @@ def _configure_session_factory(_unit_test_engine):
configure_session_factory(_unit_test_engine, expire_on_commit=False)
def persist_service_api_tenant_owner(session: Session, tenant: Tenant, owner: Account) -> TenantAccountJoin:
"""Persist the owner identity resolved by service-API app authentication.
The legacy name is retained temporarily for consumers on independent
conversion branches, but this helper no longer fabricates an execute result.
def setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_owner):
"""
membership = TenantAccountJoin(
tenant_id=tenant.id,
account_id=owner.id,
role=TenantAccountRole.OWNER,
)
owner._current_tenant = tenant
session.add_all([tenant, owner, membership])
session.commit()
return membership
Helper to stub the tenant-owner execute result for service API app authentication.
The validate_app_token decorator currently resolves the active tenant owner
via db.session.execute(select(Tenant, Account)...).one_or_none().
def persist_service_api_dataset_owner(
session: Session,
tenant: Tenant,
tenant_account_join: TenantAccountJoin,
) -> None:
"""Persist the tenant-owner mapping resolved by dataset-token authentication."""
session.add_all([tenant, tenant_account_join])
session.commit()
def setup_mock_tenant_owner_execute_result(mock_db: MagicMock, mock_tenant: object, mock_owner: object) -> None:
"""Stub the legacy owner query; SQLite-backed tests use ``persist_service_api_tenant_owner``."""
Args:
mock_db: The mocked db object
mock_tenant: Mock tenant object to return
mock_owner: Mock owner object to return from the execute result
"""
mock_db.session.execute.return_value.one_or_none.return_value = (mock_tenant, mock_owner)
def setup_mock_dataset_owner_execute_result(
mock_db: MagicMock,
mock_tenant: object,
mock_tenant_account_join: object,
) -> None:
"""Stub the legacy dataset-owner query; SQLite tests use ``persist_service_api_dataset_owner``."""
mock_db.session.execute.return_value.one_or_none.return_value = (
mock_tenant,
mock_tenant_account_join,
)
def setup_mock_dataset_owner_execute_result(mock_db, mock_tenant, mock_tenant_account_join):
"""
Helper to stub the tenant-owner execute result for dataset token authentication.
The validate_dataset_token decorator currently resolves the owner mapping via
db.session.execute(select(Tenant, TenantAccountJoin)...).one_or_none(), and
then loads the Account separately via db.session.get(...).
Args:
mock_db: The mocked db object
mock_tenant: Mock tenant object to return
mock_tenant_account_join: Mock tenant-account join object to return
"""
mock_db.session.execute.return_value.one_or_none.return_value = (mock_tenant, mock_tenant_account_join)
@@ -1,82 +1,23 @@
from types import SimpleNamespace
from typing import Any
from uuid import NAMESPACE_URL, uuid5
from unittest.mock import MagicMock
import pytest
from sqlalchemy.orm import Session
from controllers.common.agent_app_parameters import get_published_agent_app_feature_dict_and_user_input_form
from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict
from core.app.apps.agent_app.errors import AgentAppGeneratorError, AgentAppNotPublishedError
from models.agent import Agent, AgentConfigSnapshot, AgentScope, AgentSource, AgentStatus
from models.model import AppAnnotationSetting
def _stable_uuid(value: str) -> str:
return str(uuid5(NAMESPACE_URL, value))
def _app_model(*, tenant_id: str, bound_agent_id: str | None, app_model_config: object | None = None):
def _app_model(*, bound_agent_id: str | None, app_model_config=None):
return SimpleNamespace(
id=_stable_uuid(f"app:{tenant_id}"),
tenant_id=tenant_id,
id="app-1",
tenant_id="tenant-1",
bound_agent_id=bound_agent_id,
app_model_config_with_session=lambda *, session: app_model_config,
)
def _persist_agent(
session: Session,
*,
tenant_id: str,
agent_id: str,
active_config_snapshot_id: str | None,
active_config_is_published: bool,
) -> Agent:
agent = Agent(
id=agent_id,
tenant_id=tenant_id,
name="Agent",
scope=AgentScope.ROSTER,
source=AgentSource.AGENT_APP,
status=AgentStatus.ACTIVE,
active_config_snapshot_id=active_config_snapshot_id,
active_config_is_published=active_config_is_published,
)
session.add(agent)
session.commit()
return agent
def _persist_snapshot(
session: Session,
*,
snapshot_id: str,
tenant_id: str,
agent_id: str,
config_snapshot: dict[str, Any],
) -> AgentConfigSnapshot:
snapshot = AgentConfigSnapshot(
id=snapshot_id,
tenant_id=tenant_id,
agent_id=agent_id,
version=1,
config_snapshot=config_snapshot,
)
session.add(snapshot)
session.commit()
return snapshot
@pytest.mark.parametrize(
"sqlite_session",
[(Agent, AgentConfigSnapshot, AppAnnotationSetting)],
indirect=True,
)
def test_published_agent_app_parameters_use_soul_file_upload(sqlite_session: Session):
tenant_id = _stable_uuid("tenant:one")
agent_id = _stable_uuid("agent:one")
snapshot_id = _stable_uuid("snapshot:one")
def test_published_agent_app_parameters_use_soul_file_upload():
app_model_config = SimpleNamespace(
to_dict=lambda **_kwargs: {
"opening_statement": "Hi from legacy presentation config",
@@ -86,24 +27,14 @@ def test_published_agent_app_parameters_use_soul_file_upload(sqlite_session: Ses
},
}
)
app_model = _app_model(
tenant_id=tenant_id,
bound_agent_id=agent_id,
app_model_config=app_model_config,
)
_persist_agent(
sqlite_session,
tenant_id=tenant_id,
agent_id=agent_id,
active_config_snapshot_id=snapshot_id,
app_model = _app_model(bound_agent_id="agent-1", app_model_config=app_model_config)
agent = SimpleNamespace(
id="agent-1",
active_config_snapshot_id="snapshot-1",
active_config_is_published=True,
)
_persist_snapshot(
sqlite_session,
snapshot_id=snapshot_id,
tenant_id=tenant_id,
agent_id=agent_id,
config_snapshot={
snapshot = SimpleNamespace(
config_snapshot_dict={
"app_features": {
"file_upload": {
"enabled": True,
@@ -115,12 +46,14 @@ def test_published_agent_app_parameters_use_soul_file_upload(sqlite_session: Ses
}
},
"app_variables": [{"name": "topic", "type": "string", "required": True}],
},
}
)
session = MagicMock()
session.scalar.side_effect = [agent, snapshot, None]
features_dict, user_input_form = get_published_agent_app_feature_dict_and_user_input_form(
app_model,
session=sqlite_session,
session=session,
)
parameters = get_parameters_from_feature_dict(features_dict=features_dict, user_input_form=user_input_form)
@@ -136,30 +69,20 @@ def test_published_agent_app_parameters_use_soul_file_upload(sqlite_session: Ses
assert parameters["user_input_form"] == [{"text-input": {"label": "topic", "variable": "topic", "required": True}}]
@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot)], indirect=True)
def test_published_agent_app_parameters_requires_bound_agent(sqlite_session: Session):
tenant_id = _stable_uuid("tenant:unbound")
app_model = _app_model(tenant_id=tenant_id, bound_agent_id=None)
def test_published_agent_app_parameters_requires_bound_agent():
app_model = _app_model(bound_agent_id=None)
with pytest.raises(AgentAppGeneratorError, match="no bound Agent"):
get_published_agent_app_feature_dict_and_user_input_form(app_model, session=sqlite_session)
get_published_agent_app_feature_dict_and_user_input_form(app_model, session=MagicMock())
@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot)], indirect=True)
def test_published_agent_app_parameters_requires_existing_active_agent(sqlite_session: Session):
requested_tenant_id = _stable_uuid("tenant:requested")
agent_id = _stable_uuid("agent:cross-tenant")
app_model = _app_model(tenant_id=requested_tenant_id, bound_agent_id=agent_id)
_persist_agent(
sqlite_session,
tenant_id=_stable_uuid("tenant:other"),
agent_id=agent_id,
active_config_snapshot_id=None,
active_config_is_published=False,
)
def test_published_agent_app_parameters_requires_existing_active_agent():
app_model = _app_model(bound_agent_id="agent-1")
session = MagicMock()
session.scalar.return_value = None
with pytest.raises(AgentAppGeneratorError, match="no bound Agent"):
get_published_agent_app_feature_dict_and_user_input_form(app_model, session=sqlite_session)
get_published_agent_app_feature_dict_and_user_input_form(app_model, session=session)
@pytest.mark.parametrize(
@@ -169,96 +92,68 @@ def test_published_agent_app_parameters_requires_existing_active_agent(sqlite_se
False,
],
)
@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot)], indirect=True)
def test_published_agent_app_parameters_requires_published_agent(
active_config_is_published: bool, sqlite_session: Session
):
tenant_id = _stable_uuid(f"tenant:published:{active_config_is_published}")
agent_id = _stable_uuid(f"agent:published:{active_config_is_published}")
app_model = _app_model(tenant_id=tenant_id, bound_agent_id=agent_id)
_persist_agent(
sqlite_session,
tenant_id=tenant_id,
agent_id=agent_id,
def test_published_agent_app_parameters_requires_published_agent(active_config_is_published):
app_model = _app_model(bound_agent_id="agent-1")
agent = SimpleNamespace(
id="agent-1",
active_config_snapshot_id=None,
active_config_is_published=active_config_is_published,
)
session = MagicMock()
session.scalar.return_value = agent
with pytest.raises(AgentAppNotPublishedError, match="not been published"):
get_published_agent_app_feature_dict_and_user_input_form(app_model, session=sqlite_session)
get_published_agent_app_feature_dict_and_user_input_form(app_model, session=session)
@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot)], indirect=True)
def test_published_agent_app_parameters_allows_unpublished_draft_with_active_snapshot(sqlite_session: Session):
tenant_id = _stable_uuid("tenant:unpublished-draft")
agent_id = _stable_uuid("agent:unpublished-draft")
snapshot_id = _stable_uuid("snapshot:unpublished-draft")
app_model = _app_model(tenant_id=tenant_id, bound_agent_id=agent_id)
_persist_agent(
sqlite_session,
tenant_id=tenant_id,
agent_id=agent_id,
active_config_snapshot_id=snapshot_id,
def test_published_agent_app_parameters_allows_unpublished_draft_with_active_snapshot():
app_model = _app_model(bound_agent_id="agent-1")
agent = SimpleNamespace(
id="agent-1",
active_config_snapshot_id="snapshot-1",
active_config_is_published=False,
)
_persist_snapshot(
sqlite_session,
snapshot_id=snapshot_id,
tenant_id=tenant_id,
agent_id=agent_id,
config_snapshot={},
)
snapshot = SimpleNamespace(config_snapshot_dict={})
session = MagicMock()
session.scalar.side_effect = [agent, snapshot]
features_dict, user_input_form = get_published_agent_app_feature_dict_and_user_input_form(
app_model,
session=sqlite_session,
session=session,
)
assert features_dict["file_upload"]["enabled"] is True
assert user_input_form == []
@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot)], indirect=True)
def test_published_agent_app_parameters_requires_published_snapshot(sqlite_session: Session):
tenant_id = _stable_uuid("tenant:missing-snapshot")
agent_id = _stable_uuid("agent:missing-snapshot")
app_model = _app_model(tenant_id=tenant_id, bound_agent_id=agent_id)
_persist_agent(
sqlite_session,
tenant_id=tenant_id,
agent_id=agent_id,
active_config_snapshot_id=_stable_uuid("snapshot:missing"),
def test_published_agent_app_parameters_requires_published_snapshot():
app_model = _app_model(bound_agent_id="agent-1")
agent = SimpleNamespace(
id="agent-1",
active_config_snapshot_id="snapshot-1",
active_config_is_published=True,
)
session = MagicMock()
session.scalar.side_effect = [agent, None]
with pytest.raises(AgentAppGeneratorError, match="published version not found"):
get_published_agent_app_feature_dict_and_user_input_form(app_model, session=sqlite_session)
get_published_agent_app_feature_dict_and_user_input_form(app_model, session=session)
@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot)], indirect=True)
def test_published_agent_app_parameters_allows_missing_legacy_app_model_config(sqlite_session: Session):
tenant_id = _stable_uuid("tenant:no-legacy-config")
agent_id = _stable_uuid("agent:no-legacy-config")
snapshot_id = _stable_uuid("snapshot:no-legacy-config")
app_model = _app_model(tenant_id=tenant_id, bound_agent_id=agent_id)
_persist_agent(
sqlite_session,
tenant_id=tenant_id,
agent_id=agent_id,
active_config_snapshot_id=snapshot_id,
def test_published_agent_app_parameters_allows_missing_legacy_app_model_config():
app_model = _app_model(bound_agent_id="agent-1")
agent = SimpleNamespace(
id="agent-1",
active_config_snapshot_id="snapshot-1",
active_config_is_published=True,
)
_persist_snapshot(
sqlite_session,
snapshot_id=snapshot_id,
tenant_id=tenant_id,
agent_id=agent_id,
config_snapshot={},
)
snapshot = SimpleNamespace(config_snapshot_dict={})
session = MagicMock()
session.scalar.side_effect = [agent, snapshot]
features_dict, user_input_form = get_published_agent_app_feature_dict_and_user_input_form(
app_model,
session=sqlite_session,
session=session,
)
assert features_dict["file_upload"] == {
@@ -53,7 +53,6 @@ from controllers.console.app.message import (
AgentMessageFeedbackApi,
AgentMessageSuggestedQuestionApi,
)
from models.agent import AgentConfigDraftType
from services.entities.agent_entities import ComposerSaveStrategy, ComposerVariant
@@ -372,7 +371,6 @@ def test_agent_app_list_and_create_use_agent_route(
"tenant_id": "tenant-1",
"agent_id": "agent-created",
"account_id": account_id,
"draft_type": AgentConfigDraftType.DEBUG_BUILD,
"commit": False,
}
@@ -546,19 +544,8 @@ def test_agent_app_copy_uses_agent_id_and_returns_agent_detail(
}
@pytest.mark.parametrize(
("payload", "expected_draft_type"),
[
(None, AgentConfigDraftType.DEBUG_BUILD),
({"draft_type": "draft"}, AgentConfigDraftType.DRAFT),
],
)
def test_agent_debug_conversation_refresh_uses_current_user_and_draft_type(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
account_id: str,
payload: dict[str, str] | None,
expected_draft_type: AgentConfigDraftType,
def test_agent_debug_conversation_refresh_uses_current_user(
app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str
) -> None:
agent_id = "00000000-0000-0000-0000-000000000001"
captured: dict[str, object] = {}
@@ -570,9 +557,7 @@ def test_agent_debug_conversation_refresh_uses_current_user_and_draft_type(
monkeypatch.setattr(roster_controller, "_agent_roster_service", lambda *_args: FakeRosterService())
with app.test_request_context(
"/console/api/agent/00000000-0000-0000-0000-000000000001/debug-conversation/refresh",
method="POST",
json=payload,
"/console/api/agent/00000000-0000-0000-0000-000000000001/debug-conversation/refresh", method="POST"
):
response = unwrap(AgentDebugConversationRefreshApi.post)(
AgentDebugConversationRefreshApi(), MagicMock(), "tenant-1", SimpleNamespace(id=account_id), agent_id
@@ -582,12 +567,7 @@ def test_agent_debug_conversation_refresh_uses_current_user_and_draft_type(
"debug_conversation_has_messages": False,
"debug_conversation_message_count": 0,
}
assert captured == {
"tenant_id": "tenant-1",
"agent_id": agent_id,
"account_id": account_id,
"draft_type": expected_draft_type,
}
assert captured == {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id}
def test_agent_publish_and_build_draft_routes_call_composer_service(
@@ -1476,7 +1456,6 @@ def test_build_chat_finalization_helper_forces_debug_build_and_push_prompt(
"current_user": SimpleNamespace(id=account_id),
"app_model": app_model,
"agent_id": "agent-1",
"draft_type": AgentConfigDraftType.DEBUG_BUILD,
}
generate_call = cast(dict[str, object], captured["generate"])
assert generate_call["app_model"] is app_model
@@ -1541,15 +1520,11 @@ def test_agent_chat_helper_forces_agent_streaming_and_external_trace(
captured.update(kwargs)
return {"answer": "ok"}
def resolve_debug_conversation(**kwargs: object) -> str:
captured["resolve_debug_conversation"] = kwargs
return "debug-conversation-1"
monkeypatch.setattr(completion_controller.AppGenerateService, "generate", generate)
monkeypatch.setattr(
completion_controller,
"_resolve_current_user_agent_debug_conversation_id",
resolve_debug_conversation,
lambda **kwargs: "debug-conversation-1",
)
monkeypatch.setattr(
completion_controller.helper, "compact_generate_response", lambda response: {"response": response}
@@ -1569,7 +1544,6 @@ def test_agent_chat_helper_forces_agent_streaming_and_external_trace(
assert args["conversation_id"] == "debug-conversation-1"
assert args["auto_generate_name"] is False
assert args["external_trace_id"] == "trace-1"
assert cast(dict[str, object], captured["resolve_debug_conversation"])["draft_type"] == AgentConfigDraftType.DRAFT
def test_agent_chat_helper_ignores_private_exit_intent_payload_key(
@@ -1668,7 +1642,6 @@ def test_resolve_current_user_agent_debug_conversation_uses_agent_or_backing_app
current_user=SimpleNamespace(id="account-1"),
app_model=SimpleNamespace(id="app-1"),
agent_id="agent-1",
draft_type=AgentConfigDraftType.DRAFT,
)
fallback_id = completion_controller._resolve_current_user_agent_debug_conversation_id(
session="session-1", # type: ignore[arg-type]
@@ -1676,26 +1649,13 @@ def test_resolve_current_user_agent_debug_conversation_uses_agent_or_backing_app
current_user=SimpleNamespace(id="account-1"),
app_model=SimpleNamespace(id="app-1"),
agent_id=None,
draft_type=AgentConfigDraftType.DEBUG_BUILD,
)
assert explicit_id == "debug-agent-1"
assert fallback_id == "debug-backing-agent"
assert calls[1] == {
"get_or_create": {
"tenant_id": "tenant-1",
"agent_id": "agent-1",
"account_id": "account-1",
"draft_type": AgentConfigDraftType.DRAFT,
}
}
assert calls[1] == {"get_or_create": {"tenant_id": "tenant-1", "agent_id": "agent-1", "account_id": "account-1"}}
assert calls[3] == {"get_app_backing_agent": {"tenant_id": "tenant-1", "app_id": "app-1"}}
assert calls[4] == {
"get_or_create": {
"tenant_id": "tenant-1",
"agent_id": "backing-agent",
"account_id": "account-1",
"draft_type": AgentConfigDraftType.DEBUG_BUILD,
}
"get_or_create": {"tenant_id": "tenant-1", "agent_id": "backing-agent", "account_id": "account-1"}
}
@@ -1,150 +0,0 @@
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from werkzeug.exceptions import Forbidden
from controllers.console.app.error import AppNotFoundError
from controllers.console.app.wraps import agent_manage_required_for_agent_app
from core.rbac import RBACPermission, RBACResourceScope
from models.agent import AgentScope
TENANT_ID = "tenant-1"
ACCOUNT = SimpleNamespace(id="account-1")
def _guarded_view():
calls: list[dict[str, object]] = []
@agent_manage_required_for_agent_app
def view(*args, **kwargs):
calls.append(kwargs)
return "ok"
return view, calls
def _app_with_binding(binding):
app_model = MagicMock()
app_model.agent_app_binding_with_session.return_value = binding
return app_model
def _patch_guard(app_model, rbac_enabled: bool):
mock_db = MagicMock()
mock_db.session.scalar.return_value = app_model
return (
patch("controllers.console.app.wraps.db", mock_db),
patch("controllers.console.app.wraps.current_account_with_tenant", return_value=(ACCOUNT, TENANT_ID)),
patch("controllers.console.app.wraps.dify_config.RBAC_ENABLED", rbac_enabled),
)
class TestAgentManageRequiredForAgentApp:
def test_non_agent_app_passes_through_without_workspace_check(self):
view, calls = _guarded_view()
patches = _patch_guard(_app_with_binding(None), rbac_enabled=True)
with patches[0], patches[1], patches[2], patch("controllers.console.app.wraps.enforce_rbac_access") as gate:
assert view(app_id="app-1") == "ok"
gate.assert_not_called()
assert calls == [{"app_id": "app-1"}]
def test_roster_agent_app_requires_agent_manage_when_rbac_enabled(self):
view, _ = _guarded_view()
binding = SimpleNamespace(scope=AgentScope.ROSTER)
patches = _patch_guard(_app_with_binding(binding), rbac_enabled=True)
with patches[0], patches[1], patches[2], patch("controllers.console.app.wraps.enforce_rbac_access") as gate:
assert view(app_id="app-1") == "ok"
gate.assert_called_once_with(
tenant_id=TENANT_ID,
account_id=ACCOUNT.id,
resource_type=RBACResourceScope.WORKSPACE,
scene=RBACPermission.AGENT_MANAGE,
resource_required=False,
)
def test_roster_agent_app_denied_without_agent_manage(self):
view, calls = _guarded_view()
binding = SimpleNamespace(scope=AgentScope.ROSTER)
patches = _patch_guard(_app_with_binding(binding), rbac_enabled=True)
with (
patches[0],
patches[1],
patches[2],
patch("controllers.console.app.wraps.enforce_rbac_access", side_effect=Forbidden()),
):
with pytest.raises(Forbidden):
view(app_id="app-1")
assert calls == []
def test_roster_agent_app_skips_workspace_check_when_rbac_disabled(self):
view, _ = _guarded_view()
binding = SimpleNamespace(scope=AgentScope.ROSTER)
patches = _patch_guard(_app_with_binding(binding), rbac_enabled=False)
with patches[0], patches[1], patches[2], patch("controllers.console.app.wraps.enforce_rbac_access") as gate:
assert view(app_id="app-1") == "ok"
gate.assert_not_called()
def test_hidden_backing_app_is_rejected_even_without_rbac(self):
"""A workflow-only backing App is not part of the general app management plane."""
view, calls = _guarded_view()
binding = SimpleNamespace(scope=AgentScope.WORKFLOW_ONLY)
patches = _patch_guard(_app_with_binding(binding), rbac_enabled=False)
with patches[0], patches[1], patches[2]:
with pytest.raises(AppNotFoundError):
view(app_id="app-1")
assert calls == []
def test_hidden_backing_app_is_rejected_before_workspace_check(self):
view, calls = _guarded_view()
binding = SimpleNamespace(scope=AgentScope.WORKFLOW_ONLY)
patches = _patch_guard(_app_with_binding(binding), rbac_enabled=True)
with patches[0], patches[1], patches[2], patch("controllers.console.app.wraps.enforce_rbac_access") as gate:
with pytest.raises(AppNotFoundError):
view(app_id="app-1")
gate.assert_not_called()
assert calls == []
def test_binding_lookup_covers_archived_agents(self):
"""An Agent App stays gated after its roster Agent is archived."""
view, _ = _guarded_view()
app_model = _app_with_binding(SimpleNamespace(scope=AgentScope.ROSTER))
patches = _patch_guard(app_model, rbac_enabled=True)
with patches[0], patches[1], patches[2], patch("controllers.console.app.wraps.enforce_rbac_access"):
view(app_id="app-1")
_, call_kwargs = app_model.agent_app_binding_with_session.call_args
assert call_kwargs["include_archived"] is True
def test_resource_id_path_alias_is_resolved(self):
view, _ = _guarded_view()
binding = SimpleNamespace(scope=AgentScope.ROSTER)
patches = _patch_guard(_app_with_binding(binding), rbac_enabled=True)
with patches[0], patches[1], patches[2], patch("controllers.console.app.wraps.enforce_rbac_access") as gate:
assert view(resource_id="app-1") == "ok"
gate.assert_called_once()
def test_unknown_app_passes_through_for_downstream_handling(self):
view, calls = _guarded_view()
patches = _patch_guard(None, rbac_enabled=True)
with patches[0], patches[1], patches[2], patch("controllers.console.app.wraps.enforce_rbac_access") as gate:
assert view(app_id="app-1") == "ok"
gate.assert_not_called()
assert calls == [{"app_id": "app-1"}]
@@ -2,24 +2,15 @@
from __future__ import annotations
from collections.abc import Iterator
from dataclasses import dataclass
from inspect import unwrap
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from flask import Flask
from sqlalchemy import Engine, event
from sqlalchemy.orm import Session
from controllers.console.app import app_import as app_import_module
from models.account import Account
from models.base import TypeBase
from models.engine import db
from models.model import App, AppMode
from services.app_dsl_service import ImportStatus
from services.entities.dsl_entities import CheckDependenciesResult
from services.feature_service import DeploymentEdition, SystemFeatureModel, WebAppAuthModel
def _unwrap(func):
@@ -47,91 +38,17 @@ class _Result:
def _install_features(monkeypatch: pytest.MonkeyPatch, enabled: bool) -> None:
features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
webapp_auth=WebAppAuthModel(enabled=enabled),
)
features = SimpleNamespace(webapp_auth=SimpleNamespace(enabled=enabled))
monkeypatch.setattr(app_import_module.FeatureService, "get_system_features", lambda: features)
def _make_account(account_id: str = "u1") -> Account:
account = Account(name="Test User", email="test@example.com")
account.id = account_id
return account
@pytest.fixture
def app() -> Iterator[Flask]:
app = Flask(__name__)
app.config["SQLALCHEMY_DATABASE_URI"] = "sqlite:///:memory:"
db.init_app(app)
with app.app_context():
yield app
@pytest.fixture
def sqlite_app_engine(app: Flask) -> Engine:
engine = db.engine
TypeBase.metadata.create_all(engine, tables=[TypeBase.metadata.tables[App.__tablename__]])
return engine
@dataclass
class TransactionEvents:
commits: int = 0
rollbacks: int = 0
@pytest.fixture
def transaction_events() -> TransactionEvents:
"""Observe transaction decisions while keeping the controller on a real SQLAlchemy session."""
observed = TransactionEvents()
def record_commit(_session: Session) -> None:
observed.commits += 1
def record_rollback(_session: Session) -> None:
observed.rollbacks += 1
event.listen(Session, "after_commit", record_commit)
event.listen(Session, "after_rollback", record_rollback)
try:
yield observed
finally:
event.remove(Session, "after_commit", record_commit)
event.remove(Session, "after_rollback", record_rollback)
def _install_persisting_service_result(
monkeypatch: pytest.MonkeyPatch,
*,
method_name: str,
result: _Result,
) -> str:
app_id = result.app_id or "rolled-back-app"
def _return_result(import_service: app_import_module.AppDslService, *_args, **_kwargs):
import_service._session.add(
App(
id=app_id,
tenant_id="tenant-1",
name="Imported App",
mode=AppMode.WORKFLOW,
enable_site=True,
enable_api=True,
)
)
return result
monkeypatch.setattr(app_import_module.AppDslService, method_name, _return_result)
return app_id
def _assert_app_persistence(sqlite_app_engine: Engine, app_id: str, *, persisted: bool) -> None:
with Session(sqlite_app_engine) as session:
assert (session.get(App, app_id) is not None) is persisted
def _mock_session(monkeypatch: pytest.MonkeyPatch) -> MagicMock:
fake_session = MagicMock()
fake_session.__enter__.return_value = fake_session
fake_session.__exit__.return_value = None
monkeypatch.setattr(app_import_module, "db", SimpleNamespace(engine=object()))
monkeypatch.setattr(app_import_module, "Session", lambda *_args, **_kwargs: fake_session)
return fake_session
class TestAppImportApi:
@@ -140,107 +57,88 @@ class TestAppImportApi:
return app_import_module.AppImportApi()
def test_import_post_returns_failed_status_and_rolls_back(
self,
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
self, api, app: Flask, monkeypatch: pytest.MonkeyPatch
) -> None:
method = unwrap(api.post)
_install_features(monkeypatch, enabled=False)
app_id = _install_persisting_service_result(
monkeypatch,
method_name="import_app",
result=_Result(ImportStatus.FAILED, app_id=None),
session = _mock_session(monkeypatch)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.FAILED, app_id=None),
)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method(api, _make_account())
response, status = method(api, SimpleNamespace(id="u1"))
assert transaction_events.rollbacks == 1
assert transaction_events.commits == 0
_assert_app_persistence(sqlite_app_engine, app_id, persisted=False)
session.rollback.assert_called_once_with()
session.commit.assert_not_called()
assert status == 400
assert response["status"] == ImportStatus.FAILED
def test_import_post_returns_pending_status_and_commits(
self,
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
self, api, app: Flask, monkeypatch: pytest.MonkeyPatch
) -> None:
method = unwrap(api.post)
_install_features(monkeypatch, enabled=False)
app_id = _install_persisting_service_result(
monkeypatch,
method_name="import_app",
result=_Result(ImportStatus.PENDING),
session = _mock_session(monkeypatch)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.PENDING),
)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method(api, _make_account())
response, status = method(api, SimpleNamespace(id="u1"))
assert transaction_events.commits == 1
assert transaction_events.rollbacks == 0
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
session.commit.assert_called_once_with()
session.rollback.assert_not_called()
assert status == 202
assert response["status"] == ImportStatus.PENDING
def test_import_post_updates_webapp_auth_when_enabled(
self,
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
self, api, app: Flask, monkeypatch: pytest.MonkeyPatch
) -> None:
method = unwrap(api.post)
_install_features(monkeypatch, enabled=True)
app_id = _install_persisting_service_result(
monkeypatch,
method_name="import_app",
result=_Result(ImportStatus.COMPLETED, app_id="app-123"),
session = _mock_session(monkeypatch)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.COMPLETED, app_id="app-123"),
)
update_access = MagicMock()
monkeypatch.setattr(app_import_module.EnterpriseService.WebAppAuth, "update_app_access_mode", update_access)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method(api, _make_account())
response, status = method(api, SimpleNamespace(id="u1"))
assert transaction_events.commits == 1
assert transaction_events.rollbacks == 0
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
session.commit.assert_called_once_with()
session.rollback.assert_not_called()
update_access.assert_called_once_with("app-123", "private")
assert status == 200
assert response["status"] == ImportStatus.COMPLETED
def test_import_post_attaches_permission_keys_when_creating_new_app_and_rbac_enabled(
self,
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
self, api, app: Flask, monkeypatch: pytest.MonkeyPatch
) -> None:
method = _unwrap(api.post)
_install_features(monkeypatch, enabled=False)
session = _mock_session(monkeypatch)
monkeypatch.setattr(
app_import_module,
"current_account_with_tenant",
lambda: (_make_account(), "tenant-1"),
lambda: (SimpleNamespace(id="u1"), "tenant-1"),
)
monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True)
app_id = _install_persisting_service_result(
monkeypatch,
method_name="import_app",
result=_Result(ImportStatus.COMPLETED, app_id="app-123"),
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.COMPLETED, app_id="app-123"),
)
monkeypatch.setattr(
app_import_module,
@@ -251,32 +149,27 @@ class TestAppImportApi:
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
response, status = method()
assert transaction_events.commits == 1
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
session.commit.assert_called_once_with()
assert status == 200
assert response["permission_keys"] == ["app.acl.view_layout", "app.acl.edit"]
def test_import_post_does_not_attach_permission_keys_when_overwriting_existing_app(
self,
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
self, api, app: Flask, monkeypatch: pytest.MonkeyPatch
) -> None:
method = _unwrap(api.post)
_install_features(monkeypatch, enabled=False)
session = _mock_session(monkeypatch)
monkeypatch.setattr(
app_import_module,
"current_account_with_tenant",
lambda: (_make_account(), "tenant-1"),
lambda: (SimpleNamespace(id="u1"), "tenant-1"),
)
monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True)
app_id = _install_persisting_service_result(
monkeypatch,
method_name="import_app",
result=_Result(ImportStatus.COMPLETED, app_id="app-123"),
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.COMPLETED, app_id="app-123"),
)
monkeypatch.setattr(
app_import_module,
@@ -291,8 +184,7 @@ class TestAppImportApi:
):
response, status = method()
assert transaction_events.commits == 1
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
session.commit.assert_called_once_with()
assert status == 200
assert response["permission_keys"] == []
@@ -303,44 +195,35 @@ class TestAppImportConfirmApi:
return app_import_module.AppImportConfirmApi()
def test_import_confirm_returns_failed_status_and_rolls_back(
self,
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
self, api, app: Flask, monkeypatch: pytest.MonkeyPatch
) -> None:
method = unwrap(api.post)
app_id = _install_persisting_service_result(
monkeypatch,
method_name="confirm_import",
result=_Result(ImportStatus.FAILED),
session = _mock_session(monkeypatch)
monkeypatch.setattr(
app_import_module.AppDslService,
"confirm_import",
lambda *_args, **_kwargs: _Result(ImportStatus.FAILED),
)
with app.test_request_context("/console/api/apps/imports/import-1/confirm", method="POST"):
response, status = method(api, _make_account(), import_id="import-1")
response, status = method(api, SimpleNamespace(id="u1"), import_id="import-1")
assert transaction_events.rollbacks == 1
assert transaction_events.commits == 0
_assert_app_persistence(sqlite_app_engine, app_id, persisted=False)
session.rollback.assert_called_once_with()
session.commit.assert_not_called()
assert status == 400
assert response["status"] == ImportStatus.FAILED
def test_import_confirm_attaches_permission_keys_when_creating_new_app_and_rbac_enabled(
self,
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
self, api, app: Flask, monkeypatch: pytest.MonkeyPatch
) -> None:
method = _unwrap(api.post)
session = _mock_session(monkeypatch)
monkeypatch.setattr(
app_import_module,
"current_account_with_tenant",
lambda: (_make_account(), "tenant-1"),
lambda: (SimpleNamespace(id="u1"), "tenant-1"),
)
monkeypatch.setattr(
app_import_module.redis_client,
@@ -351,10 +234,10 @@ class TestAppImportConfirmApi:
),
)
monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True)
app_id = _install_persisting_service_result(
monkeypatch,
method_name="confirm_import",
result=_Result(ImportStatus.COMPLETED, app_id="app-456"),
monkeypatch.setattr(
app_import_module.AppDslService,
"confirm_import",
lambda *_args, **_kwargs: _Result(ImportStatus.COMPLETED, app_id="app-456"),
)
monkeypatch.setattr(
app_import_module,
@@ -365,25 +248,20 @@ class TestAppImportConfirmApi:
with app.test_request_context("/console/api/apps/imports/import-1/confirm", method="POST"):
response, status = method(import_id="import-1")
assert transaction_events.commits == 1
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
session.commit.assert_called_once_with()
assert status == 200
assert response["permission_keys"] == ["app.acl.view_layout", "app.acl.edit"]
def test_import_confirm_does_not_attach_permission_keys_when_overwriting_existing_app(
self,
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
self, api, app: Flask, monkeypatch: pytest.MonkeyPatch
) -> None:
method = _unwrap(api.post)
session = _mock_session(monkeypatch)
monkeypatch.setattr(
app_import_module,
"current_account_with_tenant",
lambda: (_make_account(), "tenant-1"),
lambda: (SimpleNamespace(id="u1"), "tenant-1"),
)
monkeypatch.setattr(
app_import_module.redis_client,
@@ -394,10 +272,10 @@ class TestAppImportConfirmApi:
),
)
monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True)
app_id = _install_persisting_service_result(
monkeypatch,
method_name="confirm_import",
result=_Result(ImportStatus.COMPLETED, app_id="app-456"),
monkeypatch.setattr(
app_import_module.AppDslService,
"confirm_import",
lambda *_args, **_kwargs: _Result(ImportStatus.COMPLETED, app_id="app-456"),
)
monkeypatch.setattr(
app_import_module,
@@ -408,28 +286,6 @@ class TestAppImportConfirmApi:
with app.test_request_context("/console/api/apps/imports/import-1/confirm", method="POST"):
response, status = method(import_id="import-1")
assert transaction_events.commits == 1
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
session.commit.assert_called_once_with()
assert status == 200
assert response["permission_keys"] == []
class TestAppImportCheckDependenciesApi:
def test_import_check_dependencies_returns_result(
self,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
) -> None:
api = app_import_module.AppImportCheckDependenciesApi()
method = unwrap(api.get)
monkeypatch.setattr(
app_import_module.AppDslService,
"check_dependencies",
lambda *_args, **_kwargs: CheckDependenciesResult(leaked_dependencies=[]),
)
with app.test_request_context("/console/api/apps/imports/app-1/check-dependencies", method="GET"):
response, status = method(api, app_model=App(id="app-1"))
assert status == 200
assert response["leaked_dependencies"] == []
@@ -7,28 +7,10 @@ from unittest.mock import MagicMock
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from controllers.console.app import generator as generator_module
from controllers.console.app.error import ProviderNotInitializeError
from core.errors.error import ProviderTokenNotInitError
from models.model import App, AppMode
def _persist_app(session: Session, *, tenant_id: str = "t1") -> App:
app_model = App(
id="app-1",
tenant_id=tenant_id,
name="Workflow App",
description="",
mode=AppMode.WORKFLOW,
enable_site=False,
enable_api=False,
max_active_requests=None,
)
session.add(app_model)
session.commit()
return app_model
def _model_config_payload():
@@ -84,11 +66,12 @@ def test_rule_code_generate_maps_token_error(app: Flask, monkeypatch: pytest.Mon
method(api, "t1")
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_instruction_generate_app_not_found(app: Flask, sqlite_session: Session) -> None:
def test_instruction_generate_app_not_found(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = generator_module.InstructionGenerateApi()
method = unwrap(api.post)
_persist_app(sqlite_session, tenant_id="other-tenant")
session = MagicMock()
session.scalar.return_value = None
with app.test_request_context(
"/console/api/instruction-generate",
@@ -100,21 +83,25 @@ def test_instruction_generate_app_not_found(app: Flask, sqlite_session: Session)
"model_config": _model_config_payload(),
},
):
response, status = method(api, sqlite_session, "t1")
response, status = method(api, session, "t1")
assert status == 400
assert response["error"] == "app app-1 not found"
assert sqlite_session.get(App, "app-1") is not None
stmt = session.scalar.call_args.args[0]
compiled = stmt.compile()
statement = str(compiled)
assert "apps.id" in statement
assert "apps.tenant_id" in statement
assert "app-1" in compiled.params.values()
assert "t1" in compiled.params.values()
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_instruction_generate_workflow_not_found(
app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
) -> None:
def test_instruction_generate_workflow_not_found(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = generator_module.InstructionGenerateApi()
method = unwrap(api.post)
app_model = _persist_app(sqlite_session)
app_model = SimpleNamespace(id="app-1")
session = SimpleNamespace(scalar=lambda *_args, **_kwargs: app_model)
_install_workflow_service(monkeypatch, workflow=None)
with app.test_request_context(
@@ -127,20 +114,18 @@ def test_instruction_generate_workflow_not_found(
"model_config": _model_config_payload(),
},
):
response, status = method(api, sqlite_session, "t1")
response, status = method(api, session, "t1")
assert status == 400
assert response["error"] == "workflow app-1 not found"
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_instruction_generate_node_missing(
app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
) -> None:
def test_instruction_generate_node_missing(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = generator_module.InstructionGenerateApi()
method = unwrap(api.post)
app_model = _persist_app(sqlite_session)
app_model = SimpleNamespace(id="app-1")
session = SimpleNamespace(scalar=lambda *_args, **_kwargs: app_model)
workflow = SimpleNamespace(graph_dict={"nodes": []})
_install_workflow_service(monkeypatch, workflow=workflow)
@@ -155,18 +140,18 @@ def test_instruction_generate_node_missing(
"model_config": _model_config_payload(),
},
):
response, status = method(api, sqlite_session, "t1")
response, status = method(api, session, "t1")
assert status == 400
assert response["error"] == "node node-1 not found"
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_instruction_generate_code_node(app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
def test_instruction_generate_code_node(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = generator_module.InstructionGenerateApi()
method = unwrap(api.post)
app_model = _persist_app(sqlite_session)
app_model = SimpleNamespace(id="app-1")
session = SimpleNamespace(scalar=lambda *_args, **_kwargs: app_model)
workflow = SimpleNamespace(
graph_dict={
@@ -188,19 +173,18 @@ def test_instruction_generate_code_node(app: Flask, monkeypatch: pytest.MonkeyPa
"model_config": _model_config_payload(),
},
):
response = method(api, sqlite_session, "t1")
response = method(api, session, "t1")
assert response == {"code": "x"}
assert workflow_service.app_model is app_model
assert workflow_service.session is sqlite_session
assert workflow_service.session is session
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_instruction_generate_legacy_modify(
app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
) -> None:
def test_instruction_generate_legacy_modify(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = generator_module.InstructionGenerateApi()
method = unwrap(api.post)
session = SimpleNamespace()
monkeypatch.setattr(
generator_module.LLMGenerator,
"instruction_modify_legacy",
@@ -218,15 +202,16 @@ def test_instruction_generate_legacy_modify(
"model_config": _model_config_payload(),
},
):
response = method(api, sqlite_session, "t1")
response = method(api, session, "t1")
assert response == {"instruction": "ok"}
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_instruction_generate_incompatible_params(app: Flask, sqlite_session: Session) -> None:
def test_instruction_generate_incompatible_params(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = generator_module.InstructionGenerateApi()
method = unwrap(api.post)
session = SimpleNamespace()
with app.test_request_context(
"/console/api/instruction-generate",
method="POST",
@@ -238,7 +223,7 @@ def test_instruction_generate_incompatible_params(app: Flask, sqlite_session: Se
"model_config": _model_config_payload(),
},
):
response, status = method(api, sqlite_session, "t1")
response, status = method(api, session, "t1")
assert status == 400
assert response["error"] == "incompatible parameters"
@@ -1,6 +1,5 @@
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from controllers.console.app import generator as generator_module
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
@@ -103,14 +102,12 @@ def test_structured_output_generate_exceptions(app: Flask, monkeypatch: pytest.M
method(api, "t1")
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_instruction_generate_exceptions(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
) -> None:
def test_instruction_generate_exceptions(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = generator_module.InstructionGenerateApi()
method = unwrap(api.post)
from types import SimpleNamespace
session = SimpleNamespace()
exceptions_to_test = [
(ProviderTokenNotInitError("token error"), generator_module.ProviderNotInitializeError),
@@ -138,4 +135,4 @@ def test_instruction_generate_exceptions(
},
):
with pytest.raises(expected_exception):
method(api, sqlite_session, "t1")
method(api, session, "t1")
@@ -11,7 +11,7 @@ import pytest
from flask import Flask
from sqlalchemy import func, select
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, object_session, sessionmaker
from sqlalchemy.orm import object_session, sessionmaker
from controllers.common import session as controller_session
from controllers.console.app import model_config as model_config_module
@@ -29,14 +29,11 @@ def _poison_implicit_app_config_properties(monkeypatch: pytest.MonkeyPatch) -> N
@pytest.mark.parametrize("app_mode", [AppMode.CHAT, AppMode.COMPLETION])
@pytest.mark.parametrize("sqlite_session", [(AppModelConfig,)], indirect=True)
def test_post_updates_non_agent_model_config_without_implicit_properties(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
app_mode: AppMode,
sqlite_session: Session,
) -> None:
"""Flush a non-agent config through the injected session without legacy model properties."""
api = model_config_module.ModelConfigResource()
method = unwrap(api.post)
@@ -48,16 +45,14 @@ def test_post_updates_non_agent_model_config_without_implicit_properties(
updated_at=None,
)
original_config = AppModelConfig(app_id="app-1", created_by="u1", updated_by="u1")
original_config.id = "config-0"
original_config.agent_mode = None
sqlite_session.add(original_config)
sqlite_session.commit()
_poison_implicit_app_config_properties(monkeypatch)
monkeypatch.setattr(
model_config_module.AppModelConfigService,
"validate_configuration",
lambda **_kwargs: {"pre_prompt": "hi"},
)
session = MagicMock()
def _from_model_config_dict(self, model_config):
self.pre_prompt = model_config["pre_prompt"]
@@ -67,16 +62,18 @@ def test_post_updates_non_agent_model_config_without_implicit_properties(
monkeypatch.setattr(AppModelConfig, "from_model_config_dict", _from_model_config_dict)
send_mock = MagicMock()
monkeypatch.setattr(model_config_module.app_model_config_was_updated, "send", send_mock)
session.get.return_value = original_config
with app.test_request_context("/console/api/apps/app-1/model-config", method="POST", json={"pre_prompt": "hi"}):
response = method(api, sqlite_session, "t1", "u1", app_model=app_model)
response = method(api, session, "t1", "u1", app_model=app_model)
assert send_mock.call_args.kwargs["session"] is sqlite_session
session.get.assert_called_once_with(AppModelConfig, "config-0")
session.add.assert_called_once()
session.flush.assert_called_once()
session.commit.assert_not_called()
assert send_mock.call_args.kwargs["session"] is session
assert app_model.app_model_config_id == "config-1"
assert app_model.mode == app_mode
persisted_config = sqlite_session.get(AppModelConfig, "config-1")
assert persisted_config is not None
assert persisted_config.pre_prompt == "hi"
assert response["result"] == "success"
@@ -163,11 +160,7 @@ def test_post_uses_one_session_and_rolls_back_when_signal_fails(
assert config_count == 1
@pytest.mark.parametrize("sqlite_session", [(AppModelConfig,)], indirect=True)
def test_post_encrypts_agent_tool_parameters(
app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
) -> None:
"""Agent parameter encryption reads and writes persisted model configurations."""
def test_post_encrypts_agent_tool_parameters(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
api = model_config_module.ModelConfigResource()
method = unwrap(api.post)
@@ -181,7 +174,6 @@ def test_post_encrypts_agent_tool_parameters(
_poison_implicit_app_config_properties(monkeypatch)
original_config = AppModelConfig(app_id="app-1", created_by="u1", updated_by="u1")
original_config.id = "config-0"
original_config.agent_mode = json.dumps(
{
"enabled": True,
@@ -198,8 +190,8 @@ def test_post_encrypts_agent_tool_parameters(
}
)
sqlite_session.add(original_config)
sqlite_session.commit()
session = MagicMock()
session.scalar.return_value = original_config
monkeypatch.setattr(
model_config_module.AppModelConfigService,
@@ -244,11 +236,11 @@ def test_post_encrypts_agent_tool_parameters(
monkeypatch.setattr(model_config_module.app_model_config_was_updated, "send", send_mock)
with app.test_request_context("/console/api/apps/app-1/model-config", method="POST", json={"pre_prompt": "hi"}):
response = method(api, sqlite_session, "t1", "u1", app_model=app_model)
response = method(api, session, "t1", "u1", app_model=app_model)
stored_config = sqlite_session.get(AppModelConfig, app_model.app_model_config_id)
assert stored_config is not None
stored_config = session.add.call_args[0][0]
stored_agent_mode = json.loads(stored_config.agent_mode)
session.scalar.assert_called_once()
assert app_model.mode == AppMode.AGENT_CHAT
assert stored_agent_mode["tools"][0]["tool_parameters"]["secret"] == "encrypted"
assert response["result"] == "success"
@@ -14,7 +14,6 @@ import pytest
from flask import Flask
from controllers.console.auth.activate import ActivateApi, ActivateCheckApi
from controllers.console.auth.error import InvitationAccountMismatchError
from controllers.console.error import AccountInFreezeError, AlreadyActivateError
from models.account import AccountStatus, TenantAccountRole
@@ -203,50 +202,6 @@ class TestActivateApi:
with patch("controllers.console.auth.activate.TenantService.switch_tenant") as mock:
yield mock
@patch("controllers.console.auth.activate.TenantService.create_tenant_member")
@patch("controllers.console.auth.activate.RegisterService.get_invitation_with_case_fallback")
@patch("controllers.console.auth.activate.RegisterService.revoke_token")
@patch("controllers.console.auth.activate.current_account_with_tenant")
@patch("controllers.console.auth.activate.extract_access_token", return_value="access-token")
@patch("controllers.console.auth.activate.db")
def test_activation_rejects_invitation_for_different_authenticated_account(
self,
mock_db: MagicMock,
mock_extract_access_token: MagicMock,
mock_current_account_with_tenant: MagicMock,
mock_revoke_token: MagicMock,
mock_get_invitation: MagicMock,
mock_create_tenant_member: MagicMock,
app: Flask,
mock_invitation: MagicMock,
mock_account: MagicMock,
mock_switch_tenant: MagicMock,
):
"""A logged-in account cannot consume another account's invitation token."""
current_account = MagicMock()
current_account.id = "current-account-id"
mock_account.id = "invited-account-id"
mock_account.status = AccountStatus.ACTIVE
mock_invitation["data"]["requires_setup"] = False
mock_get_invitation.return_value = mock_invitation
mock_current_account_with_tenant.return_value = (current_account, "current-workspace-id")
with app.test_request_context(
"/activate",
method="POST",
json={
"token": "valid_token",
},
):
with pytest.raises(InvitationAccountMismatchError):
ActivateApi().post()
mock_extract_access_token.assert_called_once()
mock_revoke_token.assert_not_called()
mock_create_tenant_member.assert_not_called()
mock_switch_tenant.assert_not_called()
mock_db.session.scalar.assert_not_called()
@patch("controllers.console.auth.activate.RegisterService.get_invitation_if_token_valid")
@patch("controllers.console.auth.activate.RegisterService.revoke_token")
@patch("controllers.console.auth.activate.db")
@@ -11,7 +11,7 @@ from controllers.console.auth.email_register import (
EmailRegisterResetApi,
EmailRegisterSendEmailApi,
)
from services.feature_service import DeploymentEdition, SystemFeatureModel
from services.feature_service import SystemFeatureModel
class TestEmailRegisterSendEmailApi:
@@ -34,11 +34,7 @@ class TestEmailRegisterSendEmailApi:
mock_account = MagicMock()
mock_get_account.return_value = mock_account
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
feature_flags = SystemFeatureModel(enable_email_password_login=True, is_allow_register=True)
with (
patch("controllers.console.auth.email_register.dify_config.BILLING_ENABLED", True),
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
@@ -79,11 +75,7 @@ class TestEmailRegisterCheckApi:
mock_get_data.return_value = {"email": "User@Example.com", "code": "4321"}
mock_generate_token.return_value = (None, "new-token")
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
feature_flags = SystemFeatureModel(enable_email_password_login=True, is_allow_register=True)
with (
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
@@ -131,11 +123,7 @@ class TestEmailRegisterResetApi:
mock_login.return_value = token_pair
mock_get_account.return_value = None
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
feature_flags = SystemFeatureModel(enable_email_password_login=True, is_allow_register=True)
with (
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
@@ -183,11 +171,7 @@ class TestEmailRegisterResetApi:
mock_login.return_value = token_pair
mock_get_account.return_value = None
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
feature_flags = SystemFeatureModel(enable_email_password_login=True, is_allow_register=True)
with (
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
@@ -240,11 +224,7 @@ class TestEmailRegisterResetApi:
mock_login.return_value = token_pair
mock_get_account.return_value = None
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
feature_flags = SystemFeatureModel(enable_email_password_login=True, is_allow_register=True)
with (
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
@@ -15,7 +15,7 @@ from controllers.console.auth.forgot_password import (
)
from models.account import Account
from models.engine import db
from services.feature_service import DeploymentEdition, SystemFeatureModel
from services.feature_service import SystemFeatureModel
@pytest.fixture
@@ -46,15 +46,8 @@ class TestForgotPasswordSendEmailApi:
mock_get_account.return_value = mock_account
mock_send_email.return_value = "token-123"
wraps_features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
controller_features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
is_allow_register=True,
)
wraps_features = SystemFeatureModel(enable_email_password_login=True, is_allow_register=True)
controller_features = SystemFeatureModel(is_allow_register=True)
with (
patch(
"controllers.console.auth.forgot_password.FeatureService.get_system_features",
@@ -102,10 +95,7 @@ class TestForgotPasswordCheckApi:
mock_get_data.return_value = {"email": "Admin@Example.com", "code": "4321"}
mock_generate_token.return_value = (None, "new-token")
wraps_features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
)
wraps_features = SystemFeatureModel(enable_email_password_login=True)
with (
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=wraps_features),
@@ -148,10 +138,7 @@ class TestForgotPasswordResetApi:
db.session.commit()
mock_get_account.return_value = account
wraps_features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
)
wraps_features = SystemFeatureModel(enable_email_password_login=True)
with (
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=wraps_features),
@@ -1,5 +1,5 @@
import urllib.parse
from unittest.mock import ANY, MagicMock, patch
from unittest.mock import MagicMock, patch
import pytest
from flask import Flask
@@ -91,102 +91,3 @@ def test_oauth_callback_validates_redirect_url_and_appends_new_user_flag(
assert response.headers["Location"] == (
f"{expected_target_url}{query_char}oauth_new_user={str(oauth_new_user).lower()}"
)
def test_oauth_callback_with_invitation_establishes_console_session(app: Flask) -> None:
oauth_provider = MagicMock()
oauth_provider.get_access_token.return_value = "google-access-token"
oauth_provider.get_user_info.return_value = OAuthUserInfo(
id="google-user-123",
name="Test User",
email="Invitee@Example.com",
)
account = MagicMock()
account.status = AccountStatus.ACTIVE
token_pair = MagicMock()
token_pair.access_token = "dify-access-token"
token_pair.refresh_token = "dify-refresh-token"
token_pair.csrf_token = "dify-csrf-token"
state = encode_oauth_state(invite_token="invite-token")
with (
patch("controllers.console.auth.oauth.get_oauth_providers", return_value={"google": oauth_provider}),
patch("controllers.console.auth.oauth.dify_config.CONSOLE_WEB_URL", CONSOLE_WEB_URL),
patch("controllers.console.auth.oauth.RegisterService") as register_service,
patch("controllers.console.auth.oauth.AccountService.link_account_integrate") as link_account,
patch("controllers.console.auth.oauth.AccountService.login", return_value=token_pair) as login,
patch("controllers.console.auth.oauth.TenantService.create_owner_tenant_if_not_exist") as create_workspace,
patch("controllers.console.auth.oauth.set_access_token_to_cookie") as set_access_cookie,
patch("controllers.console.auth.oauth.set_refresh_token_to_cookie") as set_refresh_cookie,
patch("controllers.console.auth.oauth.set_csrf_token_to_cookie") as set_csrf_cookie,
app.test_request_context(f"/oauth/authorize/google?code=test-code&state={state}"),
):
register_service.is_valid_invite_token.return_value = True
register_service.get_invitation_if_token_valid.return_value = {
"account": account,
"data": {
"account_id": "account-id",
"email": "invitee@example.com",
"workspace_id": "workspace-id",
},
"tenant": MagicMock(),
}
response = OAuthCallback().get("google")
assert response.status_code == 302
assert response.headers["Location"] == (f"{CONSOLE_WEB_URL}/signin/invite-settings?invite_token=invite-token")
link_account.assert_called_once_with("google", "google-user-123", account, session=ANY)
login.assert_called_once_with(account=account, session=ANY, ip_address=ANY)
create_workspace.assert_not_called()
set_access_cookie.assert_called_once_with(ANY, response, "dify-access-token")
set_refresh_cookie.assert_called_once_with(ANY, response, "dify-refresh-token")
set_csrf_cookie.assert_called_once_with(ANY, response, "dify-csrf-token")
def test_oauth_callback_with_invitation_rejects_another_account(app: Flask) -> None:
oauth_provider = MagicMock()
oauth_provider.get_access_token.return_value = "google-access-token"
oauth_provider.get_user_info.return_value = OAuthUserInfo(
id="google-user-123",
name="Test User",
email="another@example.com",
)
account = MagicMock()
account.status = AccountStatus.ACTIVE
state = encode_oauth_state(invite_token="invite-token")
with (
patch("controllers.console.auth.oauth.get_oauth_providers", return_value={"google": oauth_provider}),
patch("controllers.console.auth.oauth.dify_config.CONSOLE_WEB_URL", CONSOLE_WEB_URL),
patch("controllers.console.auth.oauth.RegisterService") as register_service,
patch("controllers.console.auth.oauth.AccountService.link_account_integrate") as link_account,
patch("controllers.console.auth.oauth.AccountService.login") as login,
patch("controllers.console.auth.oauth.set_access_token_to_cookie") as set_access_cookie,
patch("controllers.console.auth.oauth.set_refresh_token_to_cookie") as set_refresh_cookie,
patch("controllers.console.auth.oauth.set_csrf_token_to_cookie") as set_csrf_cookie,
app.test_request_context(f"/oauth/authorize/google?code=test-code&state={state}"),
):
register_service.is_valid_invite_token.return_value = True
register_service.get_invitation_if_token_valid.return_value = {
"account": account,
"data": {
"account_id": "account-id",
"email": "invitee@example.com",
"workspace_id": "workspace-id",
},
"tenant": MagicMock(),
}
response = OAuthCallback().get("google")
query = urllib.parse.parse_qs(urllib.parse.urlparse(response.headers["Location"]).query)
assert response.status_code == 302
assert query["message"] == ["This invitation was sent to another account. Please sign in with the invited account."]
assert query["invite_token"] == ["invite-token"]
link_account.assert_not_called()
login.assert_not_called()
register_service.revoke_token.assert_not_called()
set_access_cookie.assert_not_called()
set_refresh_cookie.assert_not_called()
set_csrf_cookie.assert_not_called()
@@ -1,9 +1,3 @@
"""RAG pipeline workflow controller serialization tests.
Handlers that own transactions run against real SQLite sessions so response
DTOs must be materialized before those transaction contexts close.
"""
from __future__ import annotations
from datetime import datetime
@@ -13,8 +7,6 @@ from unittest.mock import PropertyMock, patch
import pytest
from flask import Flask
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session
from controllers.console.datasets.rag_pipeline import rag_pipeline_workflow as module
from models.account import Account, TenantAccountRole
@@ -81,80 +73,90 @@ def test_draft_rag_pipeline_workflow_get_serializes_response_model(monkeypatch:
def test_published_rag_pipeline_workflows_serialize_items_before_session_closes(
app, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine
app, monkeypatch: pytest.MonkeyPatch
) -> None:
api = module.PublishedAllRagPipelineApi()
handler = unwrap_all(api.get)
session_state: dict[str, Session] = {}
session_state = {"open": False}
class _SessionContext:
def __enter__(self):
session_state["open"] = True
return object()
def __exit__(self, exc_type, exc, tb):
session_state["open"] = False
return False
class _SessionMaker:
def begin(self):
return _SessionContext()
base_workflow = _make_workflow()
class _Workflow:
def __getattr__(self, name: str):
assert session_state["session"].in_transaction() is True
assert session_state["open"] is True
return getattr(base_workflow, name)
def _get_all_published_workflow(**kwargs):
session_state["session"] = kwargs["session"]
return [_Workflow()], False
monkeypatch.setattr(module, "db", SimpleNamespace(engine=object(), session=lambda: object()))
monkeypatch.setattr(module, "sessionmaker", lambda *_args, **_kwargs: _SessionMaker())
monkeypatch.setattr(
module,
"RagPipelineService",
lambda *_args, **_kwargs: SimpleNamespace(get_all_published_workflow=_get_all_published_workflow),
lambda *_args, **_kwargs: SimpleNamespace(get_all_published_workflow=lambda **_kwargs: ([_Workflow()], False)),
)
with Session(sqlite_engine) as request_session:
monkeypatch.setattr(module, "db", SimpleNamespace(engine=sqlite_engine, session=lambda: request_session))
with app.test_request_context(
"/rag/pipelines/pipeline-1/workflows",
method="GET",
query_string={"page": 1, "limit": 10, "user_id": "", "named_only": "false"},
):
response = handler(api, _account(), pipeline=_pipeline())
with app.test_request_context(
"/rag/pipelines/pipeline-1/workflows",
method="GET",
query_string={"page": 1, "limit": 10, "user_id": "", "named_only": "false"},
):
response = handler(api, _account(), pipeline=_pipeline())
assert session_state["session"].in_transaction() is False
assert response["items"][0]["id"] == "workflow-1"
assert response["page"] == 1
assert response["limit"] == 10
assert response["has_more"] is False
def test_rag_pipeline_workflow_patch_serializes_response_model(
app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine
) -> None:
def test_rag_pipeline_workflow_patch_serializes_response_model(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
workflow = _make_workflow(marked_name="Updated release")
captured_session: dict[str, Session] = {}
def _update_workflow(**kwargs):
captured_session["session"] = kwargs["session"]
assert kwargs["session"].in_transaction() is True
return workflow
class _SessionContext:
def __enter__(self):
return object()
def __exit__(self, exc_type, exc, tb):
return False
class _SessionMaker:
def begin(self):
return _SessionContext()
monkeypatch.setattr(module, "db", SimpleNamespace(engine=object(), session=lambda: object()))
monkeypatch.setattr(module, "sessionmaker", lambda *_args, **_kwargs: _SessionMaker())
monkeypatch.setattr(
module,
"RagPipelineService",
lambda *_args, **_kwargs: SimpleNamespace(update_workflow=_update_workflow),
lambda *_args, **_kwargs: SimpleNamespace(update_workflow=lambda **_kwargs: workflow),
)
payload: dict[str, object] = {"marked_name": "Updated release"}
api = module.RagPipelineByIdApi()
handler = unwrap_all(api.patch)
with Session(sqlite_engine) as request_session:
monkeypatch.setattr(module, "db", SimpleNamespace(engine=sqlite_engine, session=lambda: request_session))
with (
app.test_request_context("/rag/pipelines/pipeline-1/workflows/workflow-1", method="PATCH", json=payload),
patch.object(type(module.console_ns), "payload", new_callable=PropertyMock, return_value=payload),
):
response = handler(
api,
_account(),
pipeline=_pipeline(),
workflow_id="workflow-1",
)
with (
app.test_request_context("/rag/pipelines/pipeline-1/workflows/workflow-1", method="PATCH", json=payload),
patch.object(type(module.console_ns), "payload", new_callable=PropertyMock, return_value=payload),
):
response = handler(
api,
_account(),
pipeline=_pipeline(),
workflow_id="workflow-1",
)
assert captured_session["session"].in_transaction() is False
assert response["id"] == "workflow-1"
assert response["marked_name"] == "Updated release"
assert response["hash"] == "hash-1"
@@ -838,29 +838,6 @@ class TestTrialChatAudioApi:
with pytest.raises(module.NoAudioUploadedError):
method(api, account, trial_app_chat)
def test_missing_file_field_returns_400(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
"""A multipart POST with no `file` field must surface as 400, not 500.
Verifies the controller passes file=None to AudioService.transcript_asr
instead of raising a KeyError that would yield HTTP 500.
"""
def fake_asr(*args, **kwargs):
assert kwargs["file"] is None
raise module.services.errors.audio.NoAudioUploadedServiceError()
api = module.TrialChatAudioApi()
method = unwrap(api.post)
with (
app.test_request_context("/", method="POST", data={}, content_type="multipart/form-data"),
patch.object(module.AudioService, "transcript_asr", side_effect=fake_asr),
):
with pytest.raises(module.NoAudioUploadedError) as exc_info:
method(api, account, trial_app_chat)
assert exc_info.value.code == 400
def test_audio_too_large(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatAudioApi()
method = unwrap(api.post)
@@ -8,8 +8,6 @@ from unittest.mock import Mock
import pytest
from flask import Flask
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, sessionmaker
from werkzeug.exceptions import HTTPException, NotFound
from controllers.console.snippets import snippet_workflow as snippet_workflow_module
@@ -38,9 +36,7 @@ def _snippet(**overrides) -> CustomizedSnippet:
@pytest.fixture(autouse=True)
def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None:
snippet_session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch) -> None:
def factory():
try:
return snippet_workflow_module.SnippetService(snippet_workflow_module._snippet_session_maker())
@@ -48,7 +44,7 @@ def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch, sqlite_engin
return snippet_workflow_module.SnippetService()
monkeypatch.setattr(snippet_workflow_module, "_snippet_service", factory)
monkeypatch.setattr(snippet_workflow_module, "_snippet_session_maker", lambda: snippet_session_maker)
monkeypatch.setattr(snippet_workflow_module, "_snippet_session_maker", Mock(return_value=Mock()))
def test_get_snippet_requires_snippet_id(app):
@@ -154,28 +150,28 @@ def test_published_workflow_get_returns_none_when_not_published(app) -> None:
assert handler(api, snippet=SimpleNamespace(id="snippet-1", is_published=False)) is None
@pytest.mark.parametrize("sqlite_session", [(CustomizedSnippet,)], indirect=True)
def test_published_workflow_post_returns_400_when_publish_fails(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
) -> None:
def test_published_workflow_post_returns_400_when_publish_fails(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
user = _account("account-1")
snippet = _snippet()
sqlite_session.add(snippet)
sqlite_session.commit()
merged_snippet = _snippet()
session = SimpleNamespace(merge=Mock(return_value=merged_snippet), commit=Mock())
def fail_publish(*, session: Session, snippet: CustomizedSnippet, account: Account):
snippet.name = "Uncommitted name"
session.add(snippet)
raise ValueError("No valid workflow found.")
class SessionContext:
def __init__(self, engine):
self.engine = engine
monkeypatch.setattr(snippet_workflow_module, "db", SimpleNamespace(engine=sqlite_engine))
def __enter__(self):
return session
def __exit__(self, exc_type, exc, tb):
return False
monkeypatch.setattr(snippet_workflow_module, "Session", SessionContext)
monkeypatch.setattr(snippet_workflow_module, "db", SimpleNamespace(engine=object()))
monkeypatch.setattr(
snippet_workflow_module,
"SnippetService",
lambda: SimpleNamespace(publish_workflow=Mock(side_effect=fail_publish)),
lambda: SimpleNamespace(publish_workflow=Mock(side_effect=ValueError("No valid workflow found."))),
)
api = snippet_workflow_module.SnippetPublishedWorkflowApi()
@@ -186,8 +182,7 @@ def test_published_workflow_post_returns_400_when_publish_fails(
assert status_code == 400
assert response == {"message": "No valid workflow found."}
sqlite_session.refresh(snippet)
assert snippet.name == "Snippet"
session.commit.assert_not_called()
def test_default_block_configs_delegates_to_service(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
@@ -208,11 +203,7 @@ def test_default_block_configs_delegates_to_service(app: Flask, monkeypatch: pyt
get_default_block_configs.assert_called_once()
def test_list_published_snippet_workflows_includes_input_fields(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
) -> None:
def test_list_published_snippet_workflows_includes_input_fields(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
workflow = SimpleNamespace(
id="workflow-1",
graph_dict={"nodes": [], "edges": []},
@@ -233,7 +224,18 @@ def test_list_published_snippet_workflows_includes_input_fields(
input_fields = [{"variable": "query", "type": "text"}]
snippet = _snippet(input_fields=json.dumps(input_fields))
monkeypatch.setattr(snippet_workflow_module, "db", SimpleNamespace(engine=sqlite_engine))
class SessionContext:
def __init__(self, engine):
self.engine = engine
def __enter__(self):
return Mock()
def __exit__(self, exc_type, exc, tb):
return False
monkeypatch.setattr(snippet_workflow_module, "Session", SessionContext)
monkeypatch.setattr(snippet_workflow_module, "db", SimpleNamespace(engine=object()))
monkeypatch.setattr(
snippet_workflow_module,
"SnippetService",
@@ -362,11 +364,8 @@ def test_restore_published_snippet_workflow_to_draft_returns_400_for_invalid_gra
assert exc.value.description == "invalid snippet workflow graph"
@pytest.mark.parametrize("sqlite_session", [(CustomizedSnippet,)], indirect=True)
def test_update_published_snippet_workflow_returns_updated_workflow(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
app: Flask, monkeypatch: pytest.MonkeyPatch
) -> None:
workflow = SimpleNamespace(
id="workflow-1",
@@ -388,15 +387,21 @@ def test_update_published_snippet_workflow_returns_updated_workflow(
user = _account("account-1")
input_fields = [{"variable": "query", "type": "text"}]
snippet = _snippet(input_fields=json.dumps(input_fields))
sqlite_session.add(snippet)
sqlite_session.commit()
session = SimpleNamespace()
update_workflow = Mock(return_value=workflow)
def update_persisted_snippet(*, session: Session, snippet: CustomizedSnippet, **_kwargs):
merged_snippet = session.merge(snippet)
merged_snippet.description = "Updated in transaction"
return workflow
class TransactionContext:
def __enter__(self):
return session
update_workflow = Mock(side_effect=update_persisted_snippet)
def __exit__(self, exc_type, exc, tb):
return False
class SessionMaker:
def begin(self):
return TransactionContext()
monkeypatch.setattr(snippet_workflow_module, "_snippet_session_maker", Mock(return_value=SessionMaker()))
monkeypatch.setattr(
snippet_workflow_module,
"SnippetService",
@@ -413,18 +418,16 @@ def test_update_published_snippet_workflow_returns_updated_workflow(
):
response = handler(api, user, snippet, workflow_id="workflow-1")
update_workflow.assert_called_once()
update_call = update_workflow.call_args.kwargs
assert isinstance(update_call["session"], Session)
assert update_call["snippet"] is snippet
assert update_call["workflow_id"] == "workflow-1"
assert update_call["account"] is user
assert update_call["data"] == {"marked_name": "v1", "marked_comment": "first version"}
update_workflow.assert_called_once_with(
session=session,
snippet=snippet,
workflow_id="workflow-1",
account=user,
data={"marked_name": "v1", "marked_comment": "first version"},
)
assert response["marked_name"] == "v1"
assert response["marked_comment"] == "first version"
assert response["input_fields"] == input_fields
sqlite_session.refresh(snippet)
assert snippet.description == "Updated in transaction"
def test_update_published_snippet_workflow_returns_400_when_no_fields(app: Flask) -> None:
@@ -438,25 +441,26 @@ def test_update_published_snippet_workflow_returns_400_when_no_fields(app: Flask
assert response == {"message": "No valid fields to update"}
@pytest.mark.parametrize("sqlite_session", [(CustomizedSnippet,)], indirect=True)
def test_update_published_snippet_workflow_raises_not_found(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
) -> None:
def test_update_published_snippet_workflow_raises_not_found(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
user = _account("account-1")
snippet = _snippet()
sqlite_session.add(snippet)
sqlite_session.commit()
def update_missing_workflow(*, session: Session, snippet: CustomizedSnippet, **_kwargs):
merged_snippet = session.merge(snippet)
merged_snippet.name = "Rolled back name"
class TransactionContext:
def __enter__(self):
return SimpleNamespace()
def __exit__(self, exc_type, exc, tb):
return False
class SessionMaker:
def begin(self):
return TransactionContext()
monkeypatch.setattr(snippet_workflow_module, "_snippet_session_maker", Mock(return_value=SessionMaker()))
monkeypatch.setattr(
snippet_workflow_module,
"SnippetService",
lambda: SimpleNamespace(update_workflow=Mock(side_effect=update_missing_workflow)),
lambda: SimpleNamespace(update_workflow=Mock(return_value=None)),
)
api = snippet_workflow_module.SnippetWorkflowByIdApi()
@@ -470,9 +474,6 @@ def test_update_published_snippet_workflow_raises_not_found(
with pytest.raises(NotFound, match="Workflow not found"):
handler(api, user, snippet, workflow_id="missing-workflow")
sqlite_session.refresh(snippet)
assert snippet.name == "Snippet"
def test_workflow_run_detail_raises_not_found_when_run_missing(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
snippet = _snippet()
@@ -1,29 +1,15 @@
from collections.abc import Iterator
from inspect import unwrap
from types import SimpleNamespace
from unittest.mock import Mock
import pytest
from flask import Flask
from sqlalchemy import event, select
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from controllers.console.snippets import snippet_workflow_draft_variable as module
from graphon.variables import StringSegment
from core.workflow.variable_prefixes import CONVERSATION_VARIABLE_NODE_ID, SYSTEM_VARIABLE_NODE_ID
from models.account import Account, AccountStatus
from models.workflow import WorkflowDraftVariable, WorkflowDraftVariableFile
from services.workflow_draft_variable_service import WorkflowDraftVariableList
pytestmark = [
pytest.mark.usefixtures("sqlite_session"),
pytest.mark.parametrize(
"sqlite_session",
[(WorkflowDraftVariable, WorkflowDraftVariableFile)],
indirect=True,
),
]
def _make_account() -> Account:
account = Account(
@@ -35,31 +21,8 @@ def _make_account() -> Account:
return account
def _make_node_variable(
variable_id: str,
*,
app_id: str = "snippet-1",
user_id: str = "user-1",
node_id: str = "llm-1",
name: str | None = None,
node_execution_id: str | None = "execution-1",
) -> WorkflowDraftVariable:
"""Create a valid node variable for persisted controller tests."""
variable = WorkflowDraftVariable.new_node_variable(
app_id=app_id,
user_id=user_id,
node_id=node_id,
name=name or variable_id,
value=StringSegment(value=f"value-{variable_id}"),
node_execution_id=node_execution_id or "execution-1",
)
variable.id = variable_id
variable.node_execution_id = node_execution_id
return variable
@pytest.fixture(autouse=True)
def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch) -> None:
def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch):
def factory():
service_factory = module.SnippetService
if isinstance(service_factory, type):
@@ -70,69 +33,33 @@ def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch) -> None:
@pytest.fixture
def app() -> Flask:
def app():
app = Flask("test_snippet_workflow_draft_variable")
app.config["TESTING"] = True
return app
@pytest.fixture
def controller_sessions(
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
) -> Iterator[scoped_session[Session]]:
"""Bind both controller session styles to the isolated SQLite engine."""
sessions = scoped_session(sessionmaker(bind=sqlite_engine, expire_on_commit=False))
monkeypatch.setattr(module, "db", SimpleNamespace(engine=sqlite_engine, session=sessions))
try:
yield sessions
finally:
sessions.remove()
def _persist_variables(sqlite_session: Session, *variables: WorkflowDraftVariable) -> None:
sqlite_session.add_all(variables)
sqlite_session.commit()
def _variable_ids(sqlite_engine: Engine) -> set[str]:
with Session(sqlite_engine) as session:
return set(session.scalars(select(WorkflowDraftVariable.id)))
def test_ensure_snippet_draft_variable_row_allowed_rejects_system_variable() -> None:
variable = WorkflowDraftVariable.new_sys_variable(
app_id="snippet-1",
user_id="user-1",
name="query",
value=StringSegment(value="query"),
node_execution_id="execution-1",
editable=True,
)
def test_ensure_snippet_draft_variable_row_allowed_rejects_system_variable():
variable = SimpleNamespace(node_id=SYSTEM_VARIABLE_NODE_ID)
with pytest.raises(module.NotFoundError, match="variable not found"):
module._ensure_snippet_draft_variable_row_allowed(variable=variable, variable_id="var-1")
def test_ensure_snippet_draft_variable_row_allowed_rejects_conversation_variable() -> None:
variable = WorkflowDraftVariable.new_conversation_variable(
app_id="snippet-1",
user_id="user-1",
name="conversation-name",
value=StringSegment(value="value"),
)
def test_ensure_snippet_draft_variable_row_allowed_rejects_conversation_variable():
variable = SimpleNamespace(node_id=CONVERSATION_VARIABLE_NODE_ID)
with pytest.raises(module.NotFoundError, match="variable not found"):
module._ensure_snippet_draft_variable_row_allowed(variable=variable, variable_id="var-1")
def test_ensure_snippet_draft_variable_row_allowed_accepts_canvas_node_variable() -> None:
variable = _make_node_variable("var-1")
def test_ensure_snippet_draft_variable_row_allowed_accepts_canvas_node_variable():
variable = SimpleNamespace(node_id="llm-1")
module._ensure_snippet_draft_variable_row_allowed(variable=variable, variable_id="var-1")
def test_conversation_variables_returns_empty_list(app: Flask) -> None:
def test_conversation_variables_returns_empty_list(app: Flask):
api = module.SnippetConversationVariableCollectionApi()
handler = unwrap(api.get)
@@ -142,7 +69,7 @@ def test_conversation_variables_returns_empty_list(app: Flask) -> None:
assert result == WorkflowDraftVariableList(variables=[])
def test_system_variables_returns_empty_list(app: Flask) -> None:
def test_system_variables_returns_empty_list(app: Flask):
api = module.SnippetSystemVariableCollectionApi()
handler = unwrap(api.get)
@@ -152,17 +79,12 @@ def test_system_variables_returns_empty_list(app: Flask) -> None:
assert result == WorkflowDraftVariableList(variables=[])
def test_delete_variable_collection_deletes_only_current_user_variables(
app: Flask,
sqlite_session: Session,
sqlite_engine: Engine,
controller_sessions: scoped_session[Session],
) -> None:
matching = _make_node_variable("matching", name="matching")
matching_second = _make_node_variable("matching-second", node_id="tool-1", name="matching-second")
other_user = _make_node_variable("other-user", user_id="user-2", name="other-user")
other_snippet = _make_node_variable("other-snippet", app_id="snippet-2", name="other-snippet")
_persist_variables(sqlite_session, matching, matching_second, other_user, other_snippet)
def test_delete_variable_collection_deletes_current_user_variables(app: Flask, monkeypatch: pytest.MonkeyPatch):
draft_var_service = SimpleNamespace(delete_user_workflow_variables=Mock())
monkeypatch.setattr(module, "WorkflowDraftVariableService", Mock(return_value=draft_var_service))
db_session = Mock()
db_session.return_value = SimpleNamespace()
monkeypatch.setattr(module.db, "session", db_session)
api = module.SnippetWorkflowVariableCollectionApi()
handler = unwrap(api.delete)
@@ -170,14 +92,11 @@ def test_delete_variable_collection_deletes_only_current_user_variables(
response = handler(api, _make_account(), snippet=SimpleNamespace(id="snippet-1"))
assert response.status_code == 204
assert _variable_ids(sqlite_engine) == {other_user.id, other_snippet.id}
assert not controller_sessions().in_transaction()
draft_var_service.delete_user_workflow_variables.assert_called_once_with("snippet-1", user_id="user-1")
db_session.commit.assert_called_once()
def test_variable_collection_get_raises_when_draft_workflow_missing(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
) -> None:
def test_variable_collection_get_raises_when_draft_workflow_missing(app: Flask, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
module,
"SnippetService",
@@ -192,37 +111,47 @@ def test_variable_collection_get_raises_when_draft_workflow_missing(
handler(api, _make_account(), snippet=SimpleNamespace(id="snippet-1"))
def test_node_variable_collection_get_lists_persisted_node_variables(
app: Flask,
sqlite_session: Session,
controller_sessions: scoped_session[Session],
) -> None:
matching = _make_node_variable("matching", name="matching")
other_node = _make_node_variable("other-node", node_id="tool-1", name="other-node")
other_user = _make_node_variable("other-user", user_id="user-2", name="other-user")
other_snippet = _make_node_variable("other-snippet", app_id="snippet-2", name="other-snippet")
_persist_variables(sqlite_session, matching, other_node, other_user, other_snippet)
def test_node_variable_collection_get_lists_node_variables(app: Flask, monkeypatch: pytest.MonkeyPatch):
variables = WorkflowDraftVariableList(variables=[SimpleNamespace(id="var-1")])
list_node_variables = Mock(return_value=variables)
class SessionContext:
def __init__(self, bind, expire_on_commit=False):
self.bind = bind
self.expire_on_commit = expire_on_commit
def __enter__(self):
return SimpleNamespace()
def __exit__(self, exc_type, exc, tb):
return False
monkeypatch.setattr(module, "Session", SessionContext)
monkeypatch.setattr(module, "db", SimpleNamespace(engine=object()))
monkeypatch.setattr(
module,
"WorkflowDraftVariableService",
Mock(return_value=SimpleNamespace(list_node_variables=list_node_variables)),
)
api = module.SnippetNodeVariableCollectionApi()
handler = unwrap(api.get)
with app.test_request_context("/"):
result = handler(api, _make_account(), snippet=SimpleNamespace(id="snippet-1"), node_id="llm-1")
assert [variable.id for variable in result.variables] == [matching.id]
assert controller_sessions().get_bind() is not None
assert result is variables
list_node_variables.assert_called_once_with("snippet-1", "llm-1", user_id="user-1")
def test_node_variable_collection_delete_deletes_only_requested_node_variables(
app: Flask,
sqlite_session: Session,
sqlite_engine: Engine,
controller_sessions: scoped_session[Session],
) -> None:
matching = _make_node_variable("matching", name="matching")
matching_second = _make_node_variable("matching-second", name="matching-second")
other_node = _make_node_variable("other-node", node_id="tool-1", name="other-node")
other_user = _make_node_variable("other-user", user_id="user-2", name="other-user")
_persist_variables(sqlite_session, matching, matching_second, other_node, other_user)
def test_node_variable_collection_delete_deletes_node_variables(app: Flask, monkeypatch: pytest.MonkeyPatch):
delete_node_variables = Mock()
draft_var_service = SimpleNamespace(delete_node_variables=delete_node_variables)
monkeypatch.setattr(module, "WorkflowDraftVariableService", Mock(return_value=draft_var_service))
db_session = Mock()
db_session.return_value = SimpleNamespace()
monkeypatch.setattr(module.db, "session", db_session)
api = module.SnippetNodeVariableCollectionApi()
handler = unwrap(api.delete)
@@ -230,102 +159,83 @@ def test_node_variable_collection_delete_deletes_only_requested_node_variables(
response = handler(api, _make_account(), snippet=SimpleNamespace(id="snippet-1"), node_id="llm-1")
assert response.status_code == 204
assert _variable_ids(sqlite_engine) == {other_node.id, other_user.id}
assert not controller_sessions().in_transaction()
delete_node_variables.assert_called_once_with("snippet-1", "llm-1", user_id="user-1")
db_session.commit.assert_called_once()
def test_variable_patch_returns_persisted_variable_without_committing_when_no_changes(
app: Flask,
sqlite_session: Session,
controller_sessions: scoped_session[Session],
) -> None:
variable = _make_node_variable("var-1")
_persist_variables(sqlite_session, variable)
session = controller_sessions()
commits: list[bool] = []
def test_variable_patch_returns_variable_when_no_changes(app: Flask, monkeypatch: pytest.MonkeyPatch):
variable = SimpleNamespace(id="var-1", app_id="snippet-1", user_id="user-1", node_id="llm-1")
draft_var_service = SimpleNamespace(get_variable=Mock(return_value=variable), update_variable=Mock())
db_session = Mock()
db_session.return_value = SimpleNamespace()
monkeypatch.setattr(module.db, "session", db_session)
monkeypatch.setattr(module, "WorkflowDraftVariableService", Mock(return_value=draft_var_service))
def record_commit(_session: Session) -> None:
commits.append(True)
event.listen(session, "after_commit", record_commit)
api = module.SnippetVariableApi()
handler = unwrap(api.patch)
try:
with app.test_request_context("/", method="PATCH", json={}):
result = handler(
api,
_make_account(),
snippet=SimpleNamespace(id="snippet-1", tenant_id="tenant-1"),
variable_id="var-1",
)
finally:
event.remove(session, "after_commit", record_commit)
assert result.id == variable.id
assert result.app_id == "snippet-1"
assert commits == []
assert session.in_transaction()
with app.test_request_context("/", method="PATCH", json={}):
result = handler(
api,
_make_account(),
snippet=SimpleNamespace(id="snippet-1", tenant_id="tenant-1"),
variable_id="var-1",
)
assert result is variable
draft_var_service.update_variable.assert_not_called()
db_session.commit.assert_not_called()
def test_variable_delete_deletes_persisted_variable(
app: Flask,
sqlite_session: Session,
sqlite_engine: Engine,
controller_sessions: scoped_session[Session],
) -> None:
variable = _make_node_variable("var-1")
retained = _make_node_variable("var-2", name="retained")
_persist_variables(sqlite_session, variable, retained)
def test_variable_delete_deletes_variable(app: Flask, monkeypatch: pytest.MonkeyPatch):
variable = SimpleNamespace(id="var-1", app_id="snippet-1", user_id="user-1", node_id="llm-1")
delete_variable = Mock()
draft_var_service = SimpleNamespace(get_variable=Mock(return_value=variable), delete_variable=delete_variable)
db_session = Mock()
db_session.return_value = SimpleNamespace()
monkeypatch.setattr(module.db, "session", db_session)
monkeypatch.setattr(module, "WorkflowDraftVariableService", Mock(return_value=draft_var_service))
api = module.SnippetVariableApi()
handler = unwrap(api.delete)
with app.test_request_context("/", method="DELETE"):
response = handler(
api,
_make_account(),
snippet=SimpleNamespace(id="snippet-1"),
variable_id=variable.id,
)
response = handler(api, _make_account(), snippet=SimpleNamespace(id="snippet-1"), variable_id="var-1")
assert response.status_code == 204
assert _variable_ids(sqlite_engine) == {retained.id}
assert not controller_sessions().in_transaction()
delete_variable.assert_called_once_with(variable)
db_session.commit.assert_called_once()
def test_variable_reset_deletes_variable_without_node_execution(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
sqlite_engine: Engine,
controller_sessions: scoped_session[Session],
) -> None:
variable = _make_node_variable("var-1", node_execution_id=None)
_persist_variables(sqlite_session, variable)
def test_variable_reset_returns_no_content_when_reset_result_is_none(app: Flask, monkeypatch: pytest.MonkeyPatch):
variable = SimpleNamespace(id="var-1", app_id="snippet-1", user_id="user-1", node_id="llm-1")
draft_workflow = SimpleNamespace(id="workflow-1")
draft_var_service = SimpleNamespace(
get_variable=Mock(return_value=variable),
reset_variable=Mock(return_value=None),
)
db_session = Mock()
db_session.return_value = SimpleNamespace()
monkeypatch.setattr(module.db, "session", db_session)
monkeypatch.setattr(module, "WorkflowDraftVariableService", Mock(return_value=draft_var_service))
monkeypatch.setattr(
module,
"SnippetService",
Mock(return_value=SimpleNamespace(get_draft_workflow=Mock(return_value=SimpleNamespace(id="workflow-1")))),
Mock(return_value=SimpleNamespace(get_draft_workflow=Mock(return_value=draft_workflow))),
)
api = module.SnippetVariableResetApi()
handler = unwrap(api.put)
with app.test_request_context("/", method="PUT"):
response = handler(
api,
_make_account(),
snippet=SimpleNamespace(id="snippet-1"),
variable_id=variable.id,
)
response = handler(api, _make_account(), snippet=SimpleNamespace(id="snippet-1"), variable_id="var-1")
assert response.status_code == 204
assert _variable_ids(sqlite_engine) == set()
assert not controller_sessions().in_transaction()
draft_var_service.reset_variable.assert_called_once_with(draft_workflow, variable)
db_session.commit.assert_called_once()
def test_environment_variables_returns_workflow_environment_variables(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
) -> None:
def test_environment_variables_returns_workflow_environment_variables(app: Flask, monkeypatch: pytest.MonkeyPatch):
env_var = SimpleNamespace(
id="env-1",
name="API_KEY",
@@ -1,11 +1,9 @@
from collections.abc import Iterator
from types import SimpleNamespace
from unittest.mock import PropertyMock, patch
from unittest.mock import MagicMock, PropertyMock, patch
import pytest
from flask import Flask
from sqlalchemy import Engine
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from sqlalchemy.orm import Session
from werkzeug.exceptions import Forbidden
import controllers.console.tag.tags as module
@@ -18,12 +16,15 @@ from controllers.console.tag.tags import (
)
from models import Account
from models.account import AccountStatus, TenantAccountRole
from models.base import TypeBase
from models.enums import TagType
from models.model import Tag
from services.tag_service import UpdateTagPayload
class SessionMatcher:
def __eq__(self, other):
return isinstance(other, Session)
def unwrap(func):
"""
Recursively unwrap decorated functions.
@@ -40,26 +41,6 @@ def app():
return app
@pytest.fixture(autouse=True)
def sqlite_db_session(
sqlite_engine: Engine,
monkeypatch: pytest.MonkeyPatch,
) -> Iterator[scoped_session[Session]]:
TypeBase.metadata.create_all(sqlite_engine, tables=[TypeBase.metadata.tables[Tag.__tablename__]])
session_registry = scoped_session(sessionmaker(bind=sqlite_engine, expire_on_commit=False))
monkeypatch.setattr(module.db, "session", session_registry)
try:
yield session_registry
finally:
session_registry.remove()
def _assert_sqlite_session(session: object, sqlite_engine: Engine) -> None:
assert isinstance(session, Session)
assert session.get_bind() is sqlite_engine
assert session.is_active
@pytest.fixture
def admin_user():
account = Account(
@@ -85,16 +66,11 @@ def readonly_user():
@pytest.fixture
def tag(sqlite_db_session: scoped_session[Session]):
tag = Tag(
tenant_id="tenant-1",
name="test-tag",
type=TagType.KNOWLEDGE,
created_by="user-1",
)
def tag():
tag = MagicMock()
tag.id = "tag-1"
sqlite_db_session.add(tag)
sqlite_db_session.commit()
tag.name = "test-tag"
tag.type = TagType.KNOWLEDGE
return tag
@@ -135,7 +111,7 @@ class TestTagListApi:
assert status == 200
assert result == [{"id": "1", "name": "tag", "type": "knowledge", "binding_count": "1"}]
def test_get_snippet_tags(self, app: Flask, sqlite_engine: Engine):
def test_get_snippet_tags(self, app: Flask):
api = TagListApi()
method = unwrap(api.get)
@@ -155,9 +131,7 @@ class TestTagListApi:
):
result, status = method(api, "tenant-1")
get_tags_mock.assert_called_once()
assert get_tags_mock.call_args.args == ("snippet", "tenant-1", None)
_assert_sqlite_session(get_tags_mock.call_args.kwargs["session"], sqlite_engine)
get_tags_mock.assert_called_once_with("snippet", "tenant-1", None, session=SessionMatcher())
assert status == 200
assert result == [{"id": "1", "name": "snippet-tag", "type": "snippet", "binding_count": "1"}]
@@ -226,7 +200,7 @@ class TestTagListApi:
class TestTagUpdateDeleteApi:
def test_patch_success(self, app: Flask, admin_user, tag, payload_patch, sqlite_engine: Engine):
def test_patch_success(self, app: Flask, admin_user, tag, payload_patch):
api = TagUpdateDeleteApi()
method = unwrap(api.patch)
@@ -250,7 +224,7 @@ class TestTagUpdateDeleteApi:
update_payload, tag_id, session = update_tags_mock.call_args.args
assert update_payload == UpdateTagPayload(name="updated")
assert tag_id == "tag-1"
_assert_sqlite_session(session, sqlite_engine)
assert session == SessionMatcher()
assert result["binding_count"] == "3"
def test_patch_forbidden(self, app: Flask, readonly_user, payload_patch):
@@ -266,7 +240,7 @@ class TestTagUpdateDeleteApi:
with pytest.raises(Forbidden):
method(api, readonly_user, "tag-1")
def test_delete_success(self, app: Flask, admin_user, sqlite_engine: Engine):
def test_delete_success(self, app: Flask, admin_user):
api = TagUpdateDeleteApi()
method = unwrap(api.delete)
@@ -276,30 +250,12 @@ class TestTagUpdateDeleteApi:
):
result, status = method(api, "tag-1")
delete_mock.assert_called_once()
tag_id, session = delete_mock.call_args.args
assert tag_id == "tag-1"
_assert_sqlite_session(session, sqlite_engine)
delete_mock.assert_called_once_with("tag-1", SessionMatcher())
assert status == 204
def test_delete_snippet_tag_checks_type_in_current_tenant(
self,
app: Flask,
admin_user,
sqlite_db_session: scoped_session[Session],
sqlite_engine: Engine,
):
def test_delete_snippet_tag_checks_type_in_current_tenant(self, app: Flask, admin_user):
api = TagUpdateDeleteApi()
method = unwrap(api.delete)
tag = Tag(
tenant_id="tenant-1",
name="snippet-tag",
type=TagType.SNIPPET,
created_by="user-1",
)
tag.id = "tag-1"
sqlite_db_session.add(tag)
sqlite_db_session.commit()
with (
app.test_request_context("/"),
@@ -308,11 +264,13 @@ class TestTagUpdateDeleteApi:
"controllers.console.tag.tags.current_account_with_tenant",
return_value=(SimpleNamespace(id="user-1"), "tenant-1"),
),
patch.object(module.db.session, "scalar", return_value=TagType.SNIPPET) as scalar_mock,
patch("controllers.console.tag.tags.enforce_rbac_access") as enforce_mock,
patch("controllers.console.tag.tags.TagService.delete_tag") as delete_mock,
):
result, status = method(api, "tag-1")
scalar_mock.assert_called_once()
enforce_mock.assert_called_once_with(
tenant_id="tenant-1",
account_id="user-1",
@@ -320,49 +278,7 @@ class TestTagUpdateDeleteApi:
scene=module.RBACPermission.SNIPPETS_CREATE_AND_MODIFY,
resource_required=False,
)
delete_mock.assert_called_once()
tag_id, session = delete_mock.call_args.args
assert tag_id == "tag-1"
_assert_sqlite_session(session, sqlite_engine)
assert result == ""
assert status == 204
def test_delete_does_not_apply_snippet_rbac_to_tag_from_another_tenant(
self,
app: Flask,
admin_user,
sqlite_db_session: scoped_session[Session],
sqlite_engine: Engine,
):
api = TagUpdateDeleteApi()
method = unwrap(api.delete)
tag = Tag(
tenant_id="other-tenant",
name="other-tenant-snippet-tag",
type=TagType.SNIPPET,
created_by="other-user",
)
tag.id = "tag-1"
sqlite_db_session.add(tag)
sqlite_db_session.commit()
with (
app.test_request_context("/"),
patch("controllers.console.tag.tags.dify_config.RBAC_ENABLED", True),
patch(
"controllers.console.tag.tags.current_account_with_tenant",
return_value=(SimpleNamespace(id="user-1"), "tenant-1"),
),
patch("controllers.console.tag.tags.enforce_rbac_access") as enforce_mock,
patch("controllers.console.tag.tags.TagService.delete_tag") as delete_mock,
):
result, status = method(api, "tag-1")
enforce_mock.assert_not_called()
delete_mock.assert_called_once()
tag_id, session = delete_mock.call_args.args
assert tag_id == "tag-1"
_assert_sqlite_session(session, sqlite_engine)
delete_mock.assert_called_once_with("tag-1", SessionMatcher())
assert result == ""
assert status == 204
@@ -3,7 +3,7 @@ from inspect import unwrap
from pytest_mock import MockerFixture
from models import Account
from services.feature_service import DeploymentEdition, FeatureModel, LimitationModel, SystemFeatureModel
from services.feature_service import FeatureModel, LimitationModel, SystemFeatureModel
def make_account() -> Account:
@@ -94,11 +94,7 @@ class TestSystemFeatureApi:
"controllers.console.feature.current_account_with_tenant_optional",
return_value=(account, "tenant-123"),
)
system_features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
is_allow_register=True,
enable_learn_app=True,
)
system_features = SystemFeatureModel(is_allow_register=True, enable_learn_app=True)
get_system_features = mocker.patch(
"controllers.console.feature.FeatureService.get_system_features",
return_value=system_features,
@@ -123,10 +119,7 @@ class TestSystemFeatureApi:
"controllers.console.feature.current_account_with_tenant_optional",
return_value=(None, None),
)
system_features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
is_allow_register=False,
)
system_features = SystemFeatureModel(is_allow_register=False)
get_system_features = mocker.patch(
"controllers.console.feature.FeatureService.get_system_features",
return_value=system_features,
@@ -1,16 +1,27 @@
"""Initialization validation tests with real setup-state persistence in SQLite."""
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import Mock
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from controllers.console import init_validate
from controllers.console.error import AlreadySetupError, InitValidateFailedError
from models.model import DifySetup
class _SessionStub:
def __init__(self, has_setup: bool):
self._has_setup = has_setup
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
def execute(self, *_args, **_kwargs):
return SimpleNamespace(scalar_one_or_none=lambda: Mock() if self._has_setup else None)
def test_get_init_status_finished(monkeypatch: pytest.MonkeyPatch) -> None:
@@ -74,15 +85,11 @@ def test_get_init_validate_status_validated_session(app: Flask, monkeypatch: pyt
assert init_validate.get_init_validate_status() is True
@pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True)
def test_get_init_validate_status_setup_exists(
app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
) -> None:
def test_get_init_validate_status_setup_exists(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(init_validate.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setenv("INIT_PASSWORD", "expected")
monkeypatch.setattr(init_validate, "db", SimpleNamespace(engine=sqlite_session.get_bind()))
sqlite_session.add(DifySetup(version="test-version"))
sqlite_session.commit()
monkeypatch.setattr(init_validate, "Session", lambda *_args, **_kwargs: _SessionStub(True))
monkeypatch.setattr(init_validate, "db", SimpleNamespace(engine=object()))
app.secret_key = "test-secret"
with app.test_request_context("/console/api/init", method="GET"):
@@ -90,13 +97,11 @@ def test_get_init_validate_status_setup_exists(
assert init_validate.get_init_validate_status() is True
@pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True)
def test_get_init_validate_status_not_validated(
app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
) -> None:
def test_get_init_validate_status_not_validated(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(init_validate.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setenv("INIT_PASSWORD", "expected")
monkeypatch.setattr(init_validate, "db", SimpleNamespace(engine=sqlite_session.get_bind()))
monkeypatch.setattr(init_validate, "Session", lambda *_args, **_kwargs: _SessionStub(False))
monkeypatch.setattr(init_validate, "db", SimpleNamespace(engine=object()))
app.secret_key = "test-secret"
with app.test_request_context("/console/api/init", method="GET"):
@@ -7,11 +7,9 @@ from unittest.mock import MagicMock
import httpx
import pytest
from flask import Flask, Response
from pydantic import SecretStr
from werkzeug.exceptions import (
BadGateway,
Forbidden,
HTTPException,
NotFound,
RequestEntityTooLarge,
ServiceUnavailable,
@@ -24,18 +22,15 @@ from controllers.console.knowledge_fs_proxy import (
_proxy_request,
_proxy_response,
proxy_knowledge_fs_get,
proxy_knowledge_fs_options,
proxy_knowledge_fs_write,
)
from controllers.console.wraps import RBACPermission
from services.knowledge_fs_operations import (
KnowledgeFSMethod,
KnowledgeFSOperation,
KnowledgeFSResponseKind,
)
from services.knowledge_fs_proxy import (
KnowledgeFSAccessDeniedError,
KnowledgeFSConfigurationError,
KnowledgeFSMethod,
KnowledgeFSOperation,
KnowledgeFSResponseKind,
KnowledgeFSRouteNotAllowedError,
KnowledgeFSUpstreamResponse,
get_knowledge_fs_operation,
@@ -51,7 +46,6 @@ def _upstream(
response: httpx.Response,
kind: KnowledgeFSResponseKind = "buffered",
*,
error_status_map: tuple[tuple[int, int], ...] = ((401, 502), (403, 403)),
max_response_bytes: int | None = None,
) -> KnowledgeFSUpstreamResponse:
operation = KnowledgeFSOperation(
@@ -61,7 +55,7 @@ def _upstream(
response_kind=kind,
required_scope="knowledge-spaces:read",
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
requires_dataset_editor=False,
max_response_bytes=max_response_bytes
or (64 * 1024 * 1024 if kind == "stream" else 25 * 1024 * 1024 if kind == "binary" else 1024 * 1024),
request_headers=(),
@@ -72,7 +66,6 @@ def _upstream(
"x-session-id",
),
response_media_types=(),
error_status_map=error_status_map,
)
return KnowledgeFSUpstreamResponse(response, kind, operation)
@@ -99,7 +92,7 @@ def _set_current_workspace(
def _bypass_policy_wrappers(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(
"controllers.console.knowledge_fs_proxy._proxy_knowledge_fs_non_get",
lambda method, path: _proxy_request(method, path),
unwrap(_proxy_knowledge_fs_non_get),
)
@@ -125,61 +118,10 @@ def test_console_blueprint_registers_generic_knowledge_fs_routes() -> None:
"/console/api/knowledge-fs/knowledge-spaces",
method="OPTIONS",
)
assert options_endpoint.endswith("proxy_knowledge_fs_options")
assert options_endpoint.endswith("proxy_knowledge_fs_get")
assert options_values == {"upstream_path": "knowledge-spaces"}
def test_proxy_options_does_not_require_an_authenticated_account(app: Flask) -> None:
with app.test_request_context(
"/console/api/knowledge-fs/knowledge-spaces",
method="OPTIONS",
headers={"Access-Control-Request-Method": "GET"},
):
response = app.make_response(proxy_knowledge_fs_options("knowledge-spaces"))
assert response.status_code == 204
def test_proxy_options_is_hidden_when_knowledge_fs_is_disabled(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr("controllers.console.knowledge_fs_proxy.dify_config.KNOWLEDGE_FS_ENABLED", False)
with app.test_request_context(
"/console/api/knowledge-fs/knowledge-spaces",
method="OPTIONS",
headers={"Access-Control-Request-Method": "GET"},
):
response = app.make_response(proxy_knowledge_fs_options("knowledge-spaces"))
assert response.status_code == 404
@pytest.mark.parametrize(
("upstream_path", "requested_method"),
[
("unregistered", "GET"),
("knowledge-spaces", "DELETE"),
("knowledge-spaces", ""),
],
)
def test_proxy_options_hides_unregistered_operations(
app: Flask,
upstream_path: str,
requested_method: str,
) -> None:
headers = {"Access-Control-Request-Method": requested_method} if requested_method else None
with app.test_request_context(
f"/console/api/knowledge-fs/{upstream_path}",
method="OPTIONS",
headers=headers,
):
response = app.make_response(proxy_knowledge_fs_options(upstream_path))
assert response.status_code == 404
def test_proxy_is_hidden_when_knowledge_fs_is_disabled(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr("controllers.console.knowledge_fs_proxy.dify_config.KNOWLEDGE_FS_ENABLED", False)
@@ -293,8 +235,10 @@ def test_read_post_applies_knowledge_rate_limit_once(
monkeypatch.setattr("controllers.console.knowledge_fs_proxy.current_account_with_tenant", current_workspace)
monkeypatch.setattr("controllers.console.wraps.current_account_with_tenant", current_workspace)
check_access = MagicMock(return_value=True)
monkeypatch.setattr("services.knowledge_fs_proxy.RBACService.CheckAccess.check", check_access)
monkeypatch.setattr(
"services.knowledge_fs_proxy.RBACService.CheckAccess.check",
MagicMock(return_value=True),
)
monkeypatch.setattr(
"controllers.console.wraps.FeatureService.get_knowledge_rate_limit",
MagicMock(return_value=MagicMock(enabled=True, limit=10)),
@@ -304,19 +248,14 @@ def test_read_post_applies_knowledge_rate_limit_once(
monkeypatch.setattr("controllers.console.wraps.redis_client.zremrangebyscore", MagicMock())
monkeypatch.setattr("controllers.console.wraps.redis_client.zcard", MagicMock(return_value=1))
proxy = MagicMock(return_value=Response(status=200))
monkeypatch.setattr("controllers.console.knowledge_fs_proxy._proxy_authorized_request", proxy)
monkeypatch.setattr("controllers.console.knowledge_fs_proxy._proxy_request", proxy)
with app.test_request_context("/console/api/knowledge-fs/knowledge-spaces", method="POST"):
response = _proxy_knowledge_fs_non_get("POST", "knowledge-spaces")
assert isinstance(response, Response)
zadd.assert_called_once()
proxy.assert_called_once()
authorization = proxy.call_args.args[0]
assert authorization.account_id == "account-1"
assert authorization.tenant_id == "tenant-1"
assert authorization.operation.operation_id == "createKnowledgeSpace"
check_access.assert_called_once()
proxy.assert_called_once_with("POST", "knowledge-spaces")
def test_denied_write_does_not_consume_the_workspace_rate_limit(
@@ -456,65 +395,6 @@ def test_generic_write_forwards_path_raw_body_and_current_tenant(
assert response.get_json()["tenantId"] == "tenant-1"
def test_generic_write_forwards_through_the_authorized_production_path(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
) -> None:
account = MagicMock(id="account-1", is_dataset_editor=True)
def current_workspace() -> tuple[MagicMock, str]:
return account, "tenant-1"
monkeypatch.setattr("controllers.console.knowledge_fs_proxy.current_account_with_tenant", current_workspace)
monkeypatch.setattr("controllers.console.wraps.current_account_with_tenant", current_workspace)
check_access = MagicMock(return_value=True)
monkeypatch.setattr("services.knowledge_fs_proxy.RBACService.CheckAccess.check", check_access)
monkeypatch.setattr(
"controllers.console.wraps.FeatureService.get_knowledge_rate_limit",
MagicMock(return_value=MagicMock(enabled=False)),
)
monkeypatch.setattr(
"services.knowledge_fs_proxy.dify_config.KNOWLEDGE_FS_BASE_URL",
"http://knowledge-fs.test",
raising=False,
)
monkeypatch.setattr(
"services.knowledge_fs_proxy.dify_config.KNOWLEDGE_FS_JWT_SECRET",
SecretStr("production-secret-with-at-least-32-bytes"),
raising=False,
)
upstream_request = MagicMock(
return_value=httpx.Response(
201,
content=b'{"id":"space-1","tenantId":"tenant-1"}',
headers={"Content-Type": "application/json"},
)
)
monkeypatch.setattr("services.knowledge_fs_proxy.ssrf_proxy.make_request", upstream_request)
route = unwrap(proxy_knowledge_fs_write)
body = b'{"idempotencyKey":"create-product-docs","name":"Product docs"}'
with app.test_request_context(
"/console/api/knowledge-fs/knowledge-spaces",
method="POST",
query_string={"source": "console"},
data=body,
content_type="application/json",
headers={"X-Trace-Id": "trace-1"},
):
response = route("knowledge-spaces")
assert isinstance(response, Response)
assert response.status_code == 201
assert response.get_json() == {"id": "space-1", "tenantId": "tenant-1"}
check_access.assert_called_once()
assert upstream_request.call_args.kwargs["method"] == "POST"
assert upstream_request.call_args.kwargs["url"] == "http://knowledge-fs.test/knowledge-spaces"
assert upstream_request.call_args.kwargs["params"] == b"source=console"
assert upstream_request.call_args.kwargs["content"] == body
assert upstream_request.call_args.kwargs["headers"]["x-trace-id"] == "trace-1"
def test_generic_write_forwards_contract_declared_request_headers(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
@@ -685,43 +565,6 @@ def test_resource_authorization_rejection_is_exposed_as_forbidden(
route("knowledge-spaces")
def test_proxy_response_applies_operation_specific_error_status_mapping() -> None:
upstream = httpx.Response(
429,
content=b'{"error":"rate limited"}',
headers={"Content-Type": "application/json"},
)
with pytest.raises(ServiceUnavailable):
_proxy_response(
_upstream(upstream, error_status_map=((429, 503),)),
tenant_id="tenant-1",
contract_response_headers=(),
max_response_bytes=1024 * 1024,
)
assert upstream.is_closed
def test_proxy_response_preserves_nonstandard_mapped_error_status() -> None:
upstream = httpx.Response(
429,
content=b'{"error":"rate limited"}',
headers={"Content-Type": "application/json"},
)
with pytest.raises(HTTPException) as exc_info:
_proxy_response(
_upstream(upstream, error_status_map=((429, 499),)),
tenant_id="tenant-1",
contract_response_headers=(),
max_response_bytes=1024 * 1024,
)
assert exc_info.value.code == 499
assert upstream.is_closed
def test_contract_response_headers_are_deduplicated_case_insensitively() -> None:
upstream = httpx.Response(
200,
@@ -776,7 +619,6 @@ def test_disallowed_non_get_route_is_hidden_as_not_found(
monkeypatch: pytest.MonkeyPatch,
method: KnowledgeFSMethod,
) -> None:
_set_current_workspace(monkeypatch)
route = unwrap(proxy_knowledge_fs_write)
with app.test_request_context("/console/api/knowledge-fs/not-a-route", method=method):
@@ -1,77 +1,11 @@
from types import SimpleNamespace
import pytest
from werkzeug.exceptions import Forbidden
from configs import dify_config
from controllers.console import workflow_run_archive
from controllers.console.workflow_run_archive import (
WorkflowRunArchiveDownloadApi,
WorkflowRunArchiveDownloadFileApi,
WorkflowRunArchiveDownloadsApi,
WorkflowRunArchivesApi,
)
from models import TenantAccountRole
@pytest.mark.parametrize(
"role",
[TenantAccountRole.EDITOR, TenantAccountRole.NORMAL, TenantAccountRole.DATASET_OPERATOR],
)
@pytest.mark.parametrize("rbac_enabled", [False, True])
def test_current_owner_or_admin_ids_rejects_non_manager(
monkeypatch: pytest.MonkeyPatch, role: TenantAccountRole, rbac_enabled: bool
) -> None:
current_user = SimpleNamespace(id="account-1", current_role=role)
monkeypatch.setattr(dify_config, "RBAC_ENABLED", rbac_enabled)
monkeypatch.setattr(
workflow_run_archive,
"current_account_with_tenant",
lambda: (current_user, "tenant-1"),
)
with pytest.raises(Forbidden):
workflow_run_archive._current_owner_or_admin_ids()
@pytest.mark.parametrize("role", [TenantAccountRole.OWNER, TenantAccountRole.ADMIN])
@pytest.mark.parametrize("rbac_enabled", [False, True])
def test_current_owner_or_admin_ids_returns_current_ids_for_manager(
monkeypatch: pytest.MonkeyPatch, role: TenantAccountRole, rbac_enabled: bool
) -> None:
current_user = SimpleNamespace(id="account-1", current_role=role)
monkeypatch.setattr(dify_config, "RBAC_ENABLED", rbac_enabled)
monkeypatch.setattr(
workflow_run_archive,
"current_account_with_tenant",
lambda: (current_user, "tenant-1"),
)
assert workflow_run_archive._current_owner_or_admin_ids() == ("tenant-1", "account-1")
@pytest.mark.parametrize(
("method", "args"),
[
(WorkflowRunArchivesApi.get, ()),
(WorkflowRunArchiveDownloadsApi.post, ()),
(WorkflowRunArchiveDownloadApi.get, ("download-1",)),
(WorkflowRunArchiveDownloadFileApi.get, ("download-1",)),
],
)
def test_workflow_run_archive_endpoints_enforce_fixed_workspace_roles(
monkeypatch: pytest.MonkeyPatch, method, args: tuple[str, ...]
) -> None:
while hasattr(method, "__wrapped__"):
method = method.__wrapped__
def reject_non_manager() -> tuple[str, str]:
raise Forbidden()
monkeypatch.setattr(workflow_run_archive, "_current_owner_or_admin_ids", reject_non_manager)
with pytest.raises(Forbidden):
method(None, *args)
@pytest.mark.parametrize(
@@ -94,4 +28,3 @@ def test_workflow_run_archive_endpoints_require_cloud_paid_plan(method) -> None:
"cloud_edition_billing_enabled",
"cloud_edition_billing_paid_plan_required",
} <= decorator_names
assert "rbac_permission_required" not in decorator_names
@@ -1,12 +1,9 @@
import inspect
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from uuid import NAMESPACE_URL, uuid5
from unittest.mock import ANY, MagicMock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from controllers.console.workspace.account import (
AccountDeleteUpdateFeedbackApi,
@@ -15,8 +12,7 @@ from controllers.console.workspace.account import (
ChangeEmailSendEmailApi,
CheckEmailUnique,
)
from models import Account, AccountIntegrate, AccountStatus, Tenant, TenantAccountJoin
from models.account import TenantAccountRole
from models import Account, AccountStatus, Tenant
from services.account_service import AccountService
from services.entities.auth_entities import (
ChangeEmailNewEmailToken,
@@ -49,38 +45,6 @@ def _build_account(email: str, account_id: str = "acc", tenant: Tenant | None =
return account
def _stable_uuid(value: str) -> str:
return str(uuid5(NAMESPACE_URL, value))
def _persist_account_with_tenant(session: Session, email: str, account_name: str = "account") -> tuple[Account, Tenant]:
tenant = Tenant(name=f"{account_name} tenant")
tenant.id = _stable_uuid(f"tenant:{account_name}")
account = Account(name=account_name, email=email, status=AccountStatus.ACTIVE)
account.id = _stable_uuid(f"account:{account_name}")
membership = TenantAccountJoin(
tenant_id=tenant.id,
account_id=account.id,
current=True,
role=TenantAccountRole.OWNER,
)
session.add_all([account, tenant, membership])
session.commit()
account._current_tenant = tenant
account.role = TenantAccountRole.OWNER
return account, tenant
@contextmanager
def _bind_database_session(session: Session):
database_session = scoped_session(sessionmaker(bind=session.get_bind(), expire_on_commit=False))
try:
with patch("extensions.ext_database.db.session", database_session):
yield database_session
finally:
database_session.remove()
def _build_change_email_token(
phase: str,
*,
@@ -423,64 +387,47 @@ class TestChangeEmailValidity:
class TestChangeEmailReset:
@patch("controllers.console.workspace.account.AccountService.send_change_email_completed_notify_email")
@patch("controllers.console.workspace.account.AccountService.update_account_email")
@patch("controllers.console.workspace.account.AccountService.revoke_change_email_token")
@patch("controllers.console.workspace.account.AccountService.get_change_email_data")
@patch("controllers.console.workspace.account.AccountService.check_email_unique")
@patch("controllers.console.workspace.account.AccountService.is_account_in_freeze")
@pytest.mark.parametrize(
"sqlite_session",
[(Account, Tenant, TenantAccountJoin, AccountIntegrate)],
indirect=True,
)
def test_should_normalize_new_email_before_update(
self,
mock_is_freeze: MagicMock,
mock_check_unique: MagicMock,
mock_get_data: MagicMock,
mock_revoke_token: MagicMock,
mock_update_account: MagicMock,
mock_send_notify: MagicMock,
app: Flask,
sqlite_session: Session,
):
current_user = _build_account("old@example.com", "acc3")
mock_is_freeze.return_value = False
mock_check_unique.return_value = True
mock_get_data.return_value = _build_change_email_token(
AccountService.CHANGE_EMAIL_PHASE_NEW_VERIFIED,
account_id="acc3",
email="new@example.com",
old_email="OLD@example.com",
)
mock_account_after_update = _build_account("new@example.com", "acc3-updated")
mock_update_account.return_value = mock_account_after_update
with _bind_database_session(sqlite_session) as database_session:
current_user, _ = _persist_account_with_tenant(
database_session(),
"old@example.com",
"email-reset-account",
)
account_integration = AccountIntegrate(
account_id=current_user.id,
provider="google",
open_id="google-user",
encrypted_token="encrypted-token",
)
database_session.add(account_integration)
database_session.commit()
mock_get_data.return_value = _build_change_email_token(
AccountService.CHANGE_EMAIL_PHASE_NEW_VERIFIED,
account_id=current_user.id,
email="new@example.com",
old_email="OLD@example.com",
)
with app.test_request_context(
"/account/change-email/reset",
method="POST",
json={"new_email": "New@Example.com", "token": "token-123"},
):
api = ChangeEmailResetApi()
method = inspect.unwrap(api.post)
method(api, current_user)
with app.test_request_context(
"/account/change-email/reset",
method="POST",
json={"new_email": "New@Example.com", "token": "token-123"},
):
api = ChangeEmailResetApi()
method = inspect.unwrap(api.post)
response = method(api, current_user)
sqlite_session.expire_all()
persisted_account = sqlite_session.get(Account, current_user.id)
assert response["email"] == "new@example.com"
assert persisted_account is not None
assert persisted_account.email == "new@example.com"
assert sqlite_session.get(AccountIntegrate, account_integration.id) is None
mock_is_freeze.assert_called_once_with("new@example.com")
mock_revoke_token.assert_called_once_with("token-123")
mock_send_notify.assert_called_once_with(email="new@example.com")
mock_is_freeze.assert_called_once_with("new@example.com")
mock_check_unique.assert_called_once_with("new@example.com", session=ANY)
mock_revoke_token.assert_called_once_with("token-123")
mock_update_account.assert_called_once_with(current_user, email="new@example.com", session=ANY)
mock_send_notify.assert_called_once_with(email="new@example.com")
@patch("controllers.console.workspace.account.AccountService.send_change_email_completed_notify_email")
@patch("controllers.console.workspace.account.AccountService.update_account_email")
@@ -716,47 +663,36 @@ class TestAccountDeletionFeedback:
class TestCheckEmailUnique:
@patch("controllers.console.workspace.account.AccountService.check_email_unique")
@patch("controllers.console.workspace.account.AccountService.is_account_in_freeze")
@pytest.mark.parametrize(
"sqlite_session",
[(Account, Tenant, TenantAccountJoin)],
indirect=True,
)
def test_should_normalize_email(
self,
mock_is_freeze: MagicMock,
app: Flask,
sqlite_session: Session,
):
def test_should_normalize_email(self, mock_is_freeze: MagicMock, mock_check_unique: MagicMock, app: Flask):
mock_is_freeze.return_value = False
mock_check_unique.return_value = True
with _bind_database_session(sqlite_session) as database_session:
_persist_account_with_tenant(database_session(), "different@test.com", "uniqueness-account")
with app.test_request_context(
"/account/change-email/check-email-unique",
method="POST",
json={"email": "Case@Test.com"},
):
api = CheckEmailUnique()
method = inspect.unwrap(api.post)
response = method(api)
with app.test_request_context(
"/account/change-email/check-email-unique",
method="POST",
json={"email": "Case@Test.com"},
):
api = CheckEmailUnique()
method = inspect.unwrap(api.post)
response = method(api)
assert response == {"result": "success"}
mock_is_freeze.assert_called_once_with("case@test.com")
mock_check_unique.assert_called_once_with("case@test.com", session=ANY)
@pytest.mark.parametrize(
"sqlite_session",
[(Account, Tenant, TenantAccountJoin)],
indirect=True,
)
def test_get_account_by_email_with_case_fallback_uses_lowercase_lookup(sqlite_session: Session):
expected_account, _ = _persist_account_with_tenant(
sqlite_session,
"mixed@test.com",
"case-fallback-account",
)
def test_get_account_by_email_with_case_fallback_uses_lowercase_lookup():
mock_session = MagicMock()
first = MagicMock()
first.scalar_one_or_none.return_value = None
second = MagicMock()
expected_account = MagicMock()
second.scalar_one_or_none.return_value = expected_account
mock_session.execute.side_effect = [first, second]
result = AccountService.get_account_by_email_with_case_fallback("Mixed@Test.com", session=sqlite_session)
result = AccountService.get_account_by_email_with_case_fallback("Mixed@Test.com", session=mock_session)
assert result is expected_account
assert mock_session.execute.call_count == 2
@@ -201,7 +201,7 @@ class TestRbacPermissionRequired:
):
assert protected_view(app_id="app-123") == "ok"
mock_extract.assert_called_once_with(RBACResourceScope.APP, "tenant-1", {"app_id": "app-123"})
mock_extract.assert_called_once_with("app", {"app_id": "app-123"})
mock_owned.assert_called_once_with("tenant-1", "account-1", "app", "app-123")
mock_check.assert_called_once_with(
"tenant-1",
@@ -307,7 +307,7 @@ class TestRbacPermissionRequired:
with app.test_request_context("/"):
request.view_args = {"app_id": "view-app"}
assert _extract_resource_id("app", "tenant-1", {"app_id": "path-app"}) == "path-app"
assert _extract_resource_id("app", {"app_id": "path-app"}) == "path-app"
def test_extract_resource_id_falls_back_to_request_view_args(self):
app = Flask(__name__)
@@ -315,59 +315,22 @@ class TestRbacPermissionRequired:
with app.test_request_context("/"):
request.view_args = {"app_id": "view-app"}
assert _extract_resource_id("app", "tenant-1") == "view-app"
assert _extract_resource_id("app") == "view-app"
def test_extract_resource_id_supports_legacy_route_aliases(self):
app = Flask(__name__)
with app.test_request_context("/apps/app-1/api-keys"):
request.view_args = {"resource_id": "app-1"}
assert _extract_resource_id(RBACResourceScope.APP, "tenant-1") == "app-1"
assert _extract_resource_id(RBACResourceScope.APP) == "app-1"
with app.test_request_context("/agent/agent-1/features"):
request.view_args = {"agent_id": "agent-1"}
assert _extract_resource_id(RBACResourceScope.APP) == "agent-1"
with app.test_request_context("/datasets/dataset-1/api-keys"):
request.view_args = {"resource_id": "dataset-1"}
assert _extract_resource_id(RBACResourceScope.DATASET, "tenant-1") == "dataset-1"
def test_extract_resource_id_resolves_agent_to_its_authz_app(self):
app = Flask(__name__)
with (
app.test_request_context("/agent/agent-1/chat-messages"),
patch("controllers.common.wraps.AgentRosterService") as mock_service,
):
request.view_args = {"agent_id": "agent-1"}
mock_service.return_value.peek_authz_app_id.return_value = "parent-app-1"
assert _extract_resource_id(RBACResourceScope.APP, "tenant-1") == "parent-app-1"
def test_extract_resource_id_scopes_agent_resolution_to_the_calling_tenant(self):
"""The tenant must reach the resolver, or an Agent id from any tenant resolves."""
app = Flask(__name__)
with (
app.test_request_context("/agent/agent-1/chat-messages"),
patch("controllers.common.wraps.AgentRosterService") as mock_service,
):
request.view_args = {"agent_id": "agent-1"}
mock_service.return_value.peek_authz_app_id.return_value = "parent-app-1"
_extract_resource_id(RBACResourceScope.APP, "tenant-9")
mock_service.return_value.peek_authz_app_id.assert_called_once_with(
tenant_id="tenant-9", agent_id="agent-1"
)
def test_extract_resource_id_keeps_agent_id_when_the_agent_does_not_resolve(self):
app = Flask(__name__)
with (
app.test_request_context("/agent/agent-1/chat-messages"),
patch("controllers.common.wraps.AgentRosterService") as mock_service,
):
request.view_args = {"agent_id": "agent-1"}
mock_service.return_value.peek_authz_app_id.return_value = None
assert _extract_resource_id(RBACResourceScope.APP, "tenant-1") == "agent-1"
assert _extract_resource_id(RBACResourceScope.DATASET) == "dataset-1"
def test_legacy_admin_decorator_noops_when_rbac_enabled(self):
@is_admin_or_owner_required

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