fix: fix conflict

This commit is contained in:
fatelei 2026-07-23 15:23:09 +08:00
commit 9237f2a14a
No known key found for this signature in database
GPG Key ID: 2F91DA05646F4EED
361 changed files with 22132 additions and 7159 deletions

View File

@ -32,12 +32,11 @@ 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.
- 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
- 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.
- 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
@ -66,7 +65,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 `vpr lint --fix --quiet` from the repository root and `pnpm -C e2e type-check`.
- Run the package-required static checks documented in `e2e/AGENTS.md`.
- Broaden verification only when the change affects hooks, tags, setup, or shared step semantics.
## Review Checklist
@ -77,6 +76,8 @@ 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.

61
.github/CODEOWNERS vendored
View File

@ -8,7 +8,6 @@
# Lint bulk suppression baselines.
/oxlint-suppressions.json
/eslint-suppressions.json
# CODEOWNERS file
/.github/CODEOWNERS @laipz8200 @crazywoola
@ -33,31 +32,9 @@
# 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
@ -111,7 +88,6 @@
/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
@ -136,11 +112,11 @@
/api/controllers/console/billing/ @hj24 @zyssyz123
# Backend - Enterprise
/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
/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
# Backend - Database Migrations
/api/migrations/ @snakevash @laipz8200 @MRZHUH
@ -153,7 +129,6 @@
# Frontend - Platform and Features
/web/config/ @lyzno1
/web/contract/ @lyzno1
/web/env.ts @lyzno1
/web/features/ @lyzno1
/web/hooks/ @lyzno1
@ -212,7 +187,6 @@
/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
@ -231,22 +205,22 @@
/web/app/components/plugins/marketplace/ @iamjoel @Yessenia-d
# Frontend - Login and Registration
/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
/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
# Frontend - Service Authentication
/web/service/base.ts @douxc @iamjoel
/web/service/base.ts @iamjoel
# Frontend - WebApp Authentication and Access Control
/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
/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
# Frontend - Explore Page
/web/app/components/explore/ @CodingOnStar @iamjoel
@ -265,7 +239,6 @@
/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

View File

@ -335,6 +335,8 @@ 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:

View File

@ -26,19 +26,30 @@ 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/**'
@ -48,8 +59,13 @@ 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/**'

View File

@ -4,9 +4,9 @@ on:
workflow_call:
inputs:
run-external-runtime:
required: false
description: Run only the prepared and external runtime suite instead of the core suites.
required: true
type: boolean
default: false
permissions:
contents: read
@ -46,6 +46,7 @@ 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
@ -54,6 +55,7 @@ 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
@ -64,7 +66,7 @@ jobs:
run: vp run e2e:full
- name: Preserve Chromium E2E report and logs
if: ${{ !cancelled() }}
if: ${{ !cancelled() && !inputs.run-external-runtime }}
run: |
if [[ -d e2e/cucumber-report ]]; then
mv e2e/cucumber-report e2e/cucumber-report-non-external
@ -74,6 +76,7 @@ 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
@ -99,7 +102,7 @@ jobs:
vp run e2e -- --tags '@browser-smoke'
- name: Preserve WebKit E2E report and logs
if: ${{ !cancelled() }}
if: ${{ !cancelled() && !inputs.run-external-runtime }}
run: |
if [[ -d e2e/cucumber-report ]]; then
mv e2e/cucumber-report e2e/cucumber-report-webkit

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](https://docs.docker.com/compose/install/) 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 v2.24.0 or later are installed on your machine:
```bash
cd dify

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, AgentStatus
from models.agent import Agent, AgentConfigDraftType, AgentStatus
from models.agent_config_entities import AgentSoulConfig
from models.enums import ApiTokenType
from models.model import ApiToken, App, IconType
@ -266,6 +266,13 @@ 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")
@ -309,6 +316,7 @@ register_schema_models(
AgentAppCopyPayload,
AgentPublishPayload,
AgentBuildDraftCheckoutPayload,
AgentDebugConversationRefreshPayload,
ComposerSavePayload,
AgentApiStatusPayload,
AgentInviteOptionsQuery,
@ -392,6 +400,7 @@ 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(
@ -439,6 +448,7 @@ 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,
@ -655,6 +665,16 @@ 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",
@ -669,10 +689,12 @@ 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,
@ -729,6 +751,7 @@ 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

View File

@ -246,7 +246,7 @@ class ModelConfigPartial(ResponseModel):
return to_timestamp(value)
class ModelConfig(ResponseModel):
class AppModelConfigResponse(ResponseModel):
opening_statement: str | None = None
suggested_questions: Any | None = Field(
default=None, validation_alias=AliasChoices("suggested_questions_list", "suggested_questions")
@ -419,7 +419,7 @@ class AppDetail(AppResponseModel):
icon_background: str | None = None
enable_site: bool
enable_api: bool
model_config_: ModelConfig | None = Field(
model_config_: AppModelConfigResponse | None = Field(
default=None,
validation_alias=AliasChoices("app_model_config", "model_config"),
alias="model_config",
@ -525,7 +525,13 @@ 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
console_ns,
RedirectUrlResponse,
SimpleResultResponse,
AppImportResponse,
AppTraceResponse,
AppModelConfigResponse,
AppDetail,
)
register_schema_models(
@ -544,10 +550,8 @@ register_schema_models(
Tag,
WorkflowPartial,
ModelConfigPartial,
ModelConfig,
AppDetailSiteResponse,
DeletedTool,
AppDetail,
AppExportResponse,
Segmentation,
PreProcessingRule,

View File

@ -49,6 +49,7 @@ 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
@ -343,14 +344,23 @@ 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
*,
session: Session,
current_tenant_id: str,
current_user: Account,
app_model: App,
agent_id: str | None,
draft_type: AgentConfigDraftType,
) -> 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))
@ -360,6 +370,7 @@ 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,
)
@ -382,6 +393,7 @@ 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.")
@ -418,6 +430,7 @@ 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,

View File

@ -359,6 +359,12 @@ class WorkflowPublishResponse(ResponseModel):
created_at: int
class SyncDraftWorkflowResponse(ResponseModel):
result: str
hash: str
updated_at: int
class WorkflowRestoreResponse(ResponseModel):
result: str
hash: str
@ -441,6 +447,7 @@ register_response_schema_models(
WorkflowOnlineUsersByApp,
WorkflowOnlineUsersResponse,
WorkflowPublishResponse,
SyncDraftWorkflowResponse,
WorkflowRestoreResponse,
DefaultBlockConfigsResponse,
DefaultBlockConfigResponse,
@ -556,14 +563,7 @@ class DraftWorkflowApi(Resource):
@console_ns.response(
200,
"Draft workflow synced successfully",
console_ns.model(
"SyncDraftWorkflowResponse",
{
"result": fields.String,
"hash": fields.String,
"updated_at": fields.String,
},
),
console_ns.models[SyncDraftWorkflowResponse.__name__],
)
@console_ns.response(400, "Invalid workflow configuration")
@console_ns.response(403, "Permission denied")
@ -618,11 +618,14 @@ class DraftWorkflowApi(Resource):
except VariableError as e:
raise InvalidArgumentError(description=str(e))
return {
"result": "success",
"hash": workflow.unique_hash,
"updated_at": TimestampField().format(workflow.updated_at or workflow.created_at),
}
return dump_response(
SyncDraftWorkflowResponse,
{
"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")

View File

@ -7,10 +7,13 @@ 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
@ -136,6 +139,12 @@ 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
@ -146,6 +155,11 @@ 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()

View File

@ -13,6 +13,12 @@ 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."

View File

@ -6,6 +6,7 @@ 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
@ -127,6 +128,20 @@ 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")
@ -195,16 +210,26 @@ 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_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.")
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}")
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin/invite-settings?invite_token={invite_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)
try:
account, oauth_new_user = _generate_account(provider, user_info, timezone=timezone, language=language)
@ -239,21 +264,10 @@ 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()}"
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
return _redirect_with_console_session(account, target_url)
def _get_account_by_openid_or_email(provider: str, user_info: OAuthUserInfo) -> Account | None:

View File

@ -8,8 +8,9 @@ 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. Upstream 401 responses become 502 so they cannot
trigger Dify browser-session recovery; resource-level 403 responses remain 403.
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.
"""
from __future__ import annotations
@ -27,9 +28,11 @@ from werkzeug.exceptions import (
BadGateway,
Forbidden,
GatewayTimeout,
HTTPException,
NotFound,
RequestEntityTooLarge,
ServiceUnavailable,
default_exceptions,
)
from configs import dify_config
@ -41,21 +44,28 @@ 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",
@ -128,27 +138,25 @@ def _translate_proxy_error(exc: Exception, *, tenant_id: str) -> NoReturn:
def _knowledge_fs_operation_access_required(
view: Callable[[KnowledgeFSMethod, str], ResponseReturnValue],
view: Callable[[KnowledgeFSAuthorization], ResponseReturnValue],
) -> Callable[[KnowledgeFSMethod, str], ResponseReturnValue]:
"""Authorize one declared operation before billing and request-body work."""
@wraps(view)
def decorated(method: KnowledgeFSMethod, upstream_path: str) -> ResponseReturnValue:
try:
operation = get_knowledge_fs_operation(method, upstream_path)
except KnowledgeFSRouteNotAllowedError as exc:
raise NotFound() from exc
current_user, tenant_id = current_account_with_tenant()
try:
authorize_knowledge_fs_request(
authorization = authorize_knowledge_fs_request(
account=current_user,
tenant_id=tenant_id,
operation=operation,
method=method,
path=upstream_path,
)
except KnowledgeFSRouteNotAllowedError as exc:
raise NotFound() from exc
except KnowledgeFSAccessDeniedError as exc:
_translate_proxy_error(exc, tenant_id=tenant_id)
return view(method, upstream_path)
return view(authorization)
return decorated
@ -190,21 +198,26 @@ def _proxy_response(
"""Expose raw content, status, and allowlisted headers from KnowledgeFS.
Raises:
BadGateway: KnowledgeFS rejects the configured server credential.
Forbidden: KnowledgeFS denies the account access to the requested resource.
HTTPException: KnowledgeFS returns a status normalized by the operation contract.
"""
upstream = upstream_result.response
if upstream.status_code == HTTPStatus.UNAUTHORIZED:
mapped_status = dict(upstream_result.operation.error_status_map).get(upstream.status_code)
if mapped_status is not None:
upstream.close()
logger.error(
"KnowledgeFS rejected the Dify server credential with HTTP %s for tenant_id=%s",
upstream.status_code,
tenant_id,
)
raise BadGateway("KnowledgeFS authentication failed")
if upstream.status_code == HTTPStatus.FORBIDDEN:
upstream.close()
raise Forbidden()
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)
allowed_header_names = dict.fromkeys(
name.lower() for name in (*_RESPONSE_HEADER_ALLOWLIST, *contract_response_headers)
@ -237,26 +250,21 @@ def _proxy_response(
return Response(content, status=upstream.status_code, headers=headers)
def _proxy_request(method: KnowledgeFSMethod, upstream_path: str) -> Response:
"""Forward the current raw request and return its filtered upstream response.
The call performs one outbound KnowledgeFS request. Integration failures are
converted to Console HTTP exceptions for the outer JSON error adapter.
"""
def _proxy_current_request(
*,
method: KnowledgeFSMethod,
tenant_id: str,
forward: _KnowledgeFSRequestForwarder,
) -> Response:
"""Forward the current raw request through one preconfigured service entry."""
if not dify_config.KNOWLEDGE_FS_ENABLED:
raise NotFound()
current_user, tenant_id = current_account_with_tenant()
try:
proxy_result = proxy_knowledge_fs_request(
account=current_user,
method=method,
path=upstream_path,
tenant_id=tenant_id,
accept=request.headers.get("Accept"),
content_type=request.content_type,
query=request.query_string or None,
body=_request_body() if method != "GET" else None,
request_headers=request.headers,
proxy_result = forward(
request.headers.get("Accept"),
request.content_type,
request.query_string or None,
_request_body() if method != "GET" else None,
)
except (
KnowledgeFSConfigurationError,
@ -274,20 +282,99 @@ def _proxy_request(method: KnowledgeFSMethod, upstream_path: str) -> Response:
)
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(
method: KnowledgeFSMethod,
upstream_path: str,
authorization: KnowledgeFSAuthorization,
) -> ResponseReturnValue:
"""Apply knowledge billing checks to one allowlisted non-GET operation."""
return _proxy_request(method, upstream_path)
return _proxy_authorized_request(authorization)
@bp.route(
"/knowledge-fs/<path:upstream_path>",
methods=["GET", "OPTIONS"],
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"],
provide_automatic_options=False,
)
@_console_api_errors

View File

@ -67,6 +67,15 @@ 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
@ -645,6 +654,7 @@ 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

View File

@ -1,6 +1,6 @@
from typing import Any, Self
from pydantic import AliasChoices, Field, computed_field
from pydantic import AliasChoices, Field
from sqlalchemy import select
from werkzeug.exceptions import Forbidden
@ -9,11 +9,13 @@ 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, Site
from models.model import App, EndUser, IconType, Site
from services.feature_service import FeatureModel, FeatureService
from services.file_service import FileService
class WebSiteResponse(ResponseModel):
@ -32,11 +34,7 @@ class WebSiteResponse(ResponseModel):
prompt_public: bool | None = None
show_workflow_steps: bool | None = None
use_icon_as_answer_icon: bool | 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)
icon_url: str | None = None
class WebModelConfigResponse(ResponseModel):
@ -88,6 +86,7 @@ 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:
@ -102,6 +101,7 @@ 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,6 +123,15 @@ 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")
@ -159,4 +168,5 @@ 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")

View File

@ -682,42 +682,27 @@ 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 == effective_draft_type,
AgentConfigDraft.draft_type == AgentConfigDraftType.DEBUG_BUILD,
AgentConfigDraft.account_id == account_id,
)
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
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
raise AgentAppGeneratorError("Agent build draft not found")
@staticmethod
def _resolve_agent_by_id(

View File

@ -1,7 +1,7 @@
import inspect
import json
import logging
from collections.abc import Callable, Generator
from collections.abc import Callable, Generator, Mapping
from typing import Any, cast
from urllib.parse import unquote
@ -23,6 +23,7 @@ from core.plugin.impl.exc import (
PluginLLMPollingUnsupportedError,
PluginNotFoundError,
PluginPermissionDeniedError,
PluginRuntimeError,
PluginUniqueIdentifierError,
)
from core.trigger.errors import (
@ -375,6 +376,18 @@ 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__:

View File

@ -49,6 +49,18 @@ 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"

View File

@ -58,7 +58,6 @@ 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(
@ -87,14 +86,14 @@ class Tool(ABC):
return result
def _transform_tool_parameters_type(self, tool_parameters: dict[str, Any]) -> dict[str, Any]:
"""
Transform tool parameters type
"""
# Temp fix for the issue that the tool parameters will be converted to empty while validating the credentials
"""Transform declared tool parameter values without resolving runtime schemas."""
result = deepcopy(tool_parameters)
for parameter in self.entity.parameters or []:
if parameter.name in tool_parameters:
result[parameter.name] = parameter.type.cast_value(tool_parameters[parameter.name])
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])
return result
@ -196,17 +195,31 @@ class Tool(ABC):
}:
continue
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)
)
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.setdefault("description", parameter.llm_description or "")
if parameter.type == ToolParameter.ToolParameterType.SELECT and parameter.options:
if (
not is_multiple_select
and parameter.type == ToolParameter.ToolParameterType.SELECT
and parameter.options
):
parameter_schema["enum"] = [option.value for option in parameter.options]
schema["properties"][parameter.name] = parameter_schema

View File

@ -292,9 +292,7 @@ class ToolInvokeMessageBinary(BaseModel):
class ToolParameter(PluginParameter):
"""
Overrides type
"""
"""Tool-specific parameter declaration and invocation-value normalization."""
class ToolParameterType(StrEnum):
"""
@ -333,12 +331,28 @@ 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,
@ -378,8 +392,25 @@ class ToolParameter(PluginParameter):
options=option_objs,
)
def init_frontend_parameter(self, value: Any):
return init_frontend_parameter(self, self.type, value)
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
class ToolProviderIdentity(BaseModel):

View File

@ -12,6 +12,7 @@ import json
import subprocess
import sys
import tempfile
from copy import deepcopy
from pathlib import Path
from typing import Any, Literal, TypedDict
@ -24,6 +25,16 @@ 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):
@ -38,6 +49,7 @@ 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[
@ -70,6 +82,7 @@ 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()
@ -101,7 +114,15 @@ def main() -> None:
)
document: dict[str, Any] = json.loads(openapi_content)
validate_declarations(document, console_contract_declarations())
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")
if args.update_lock:
LOCK_PATH.write_text(
@ -147,8 +168,7 @@ 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: ContractDeclaration = {
"operation_id": operation_id,
expected: dict[DeclarationField, object] = {
"method": method.upper(),
"path": path[1:],
"required_scope": required_scope(operation),
@ -166,11 +186,114 @@ 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_proxy import KNOWLEDGE_FS_CONSOLE_OPERATIONS
from services.knowledge_fs_operations import KNOWLEDGE_FS_CONSOLE_OPERATIONS
return tuple(
{
@ -183,11 +306,18 @@ 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:

View File

@ -119,6 +119,19 @@ 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)

View File

@ -92,3 +92,21 @@ 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,
)

View File

@ -31,6 +31,16 @@ 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.

View File

@ -1,5 +1,5 @@
{
"commit": "4310e2d582d25e7de58183f27720afab01e123cf",
"openapiSha256": "5827ca930ce38462bfd1b2bef387efbf37eb7ffcaedde4558af2fbaeccbfbc4b",
"commit": "a0f50470612cc0b3656f89e4f2435aaf412e6e3b",
"openapiSha256": "f18910e9c45a64f0855e0643a7a626fb2889021b4f943458de86c6bd2469facb",
"repository": "https://github.com/langgenius/knowledge-fs"
}

View File

@ -9,6 +9,8 @@ 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
@ -100,6 +102,20 @@ 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)
@ -121,6 +137,7 @@ 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)

View File

@ -221,8 +221,11 @@ def current_timestamp() -> int:
def email(email):
# Define a regex pattern for email addresses
pattern = r"^[\w\.!#$%&'*+\-/=?^_`{|}~]+@([\w-]+\.)+[\w-]{2,}$"
# Check if the email matches the pattern
if re.match(pattern, email) is not None:
# 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:
return email
error = f"{email} is not a valid email."

View File

@ -0,0 +1,77 @@
"""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")

View File

@ -222,11 +222,13 @@ class Agent(DefaultFieldsMixin, Base):
class AgentDebugConversation(DefaultFieldsMixin, Base):
"""Per-account console debug conversation for an Agent App.
"""Per-account, per-draft 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 pointer used by console debug chat.
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.
"""
__tablename__ = "agent_debug_conversations"
@ -236,7 +238,8 @@ class AgentDebugConversation(DefaultFieldsMixin, Base):
"tenant_id",
"agent_id",
"account_id",
name="agent_debug_conversation_agent_account_unique",
"draft_type",
name="agent_debug_conversation_agent_account_draft_type_unique",
),
Index("agent_debug_conversation_conversation_idx", "conversation_id"),
Index("agent_debug_conversation_account_idx", "tenant_id", "account_id"),
@ -246,6 +249,12 @@ 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)

View File

@ -44,18 +44,28 @@ _DECLARED_OUTPUT_CHILDREN_JSON_SCHEMA = {
},
"description": {"anyOf": [{"type": "string"}, {"type": "null"}]},
"required": {"type": "boolean"},
"file": {"type": "object", "additionalProperties": True},
"file": {
"anyOf": [
{"type": "object", "additionalProperties": True},
{"type": "null"},
]
},
"array_item": {
"type": "object",
"additionalProperties": True,
"properties": {
"type": {
"type": "string",
"enum": [item.value for item in DeclaredOutputType],
"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}},
},
},
"description": {"anyOf": [{"type": "string"}, {"type": "null"}]},
"children": {"type": "array", "items": {"type": "object", "additionalProperties": True}},
},
{"type": "null"},
]
},
"children": {"type": "array", "items": {"type": "object", "additionalProperties": True}},
},

View File

@ -260,7 +260,12 @@ 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
@ -531,6 +536,7 @@ 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
@ -958,6 +964,12 @@ 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 |
@ -11332,6 +11344,12 @@ 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 |
@ -13278,7 +13296,7 @@ Model class for AI model.
| maintainer | string | | No |
| max_active_requests | integer | | No |
| mode | string | | Yes |
| model_config | [ModelConfig](#modelconfig) | | No |
| model_config | [AppModelConfigResponse](#appmodelconfigresponse) | | No |
| name | string | | Yes |
| permission_keys | [ string ] | | No |
| role | string | | No |
@ -13899,6 +13917,12 @@ 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 |
@ -15348,7 +15372,6 @@ 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 |
@ -15358,7 +15381,8 @@ 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_compatible_with_agent | string | | Yes |
| mode | string | | Yes |
| model_config | [AppModelConfigResponse](#appmodelconfigresponse) | | No |
| name | string | | Yes |
| permission_keys | [ string ] | | No |
| tags | [ [Tag](#tag) ] | | No |
@ -15420,7 +15444,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 | [ModelConfig](#modelconfig) | | No |
| model_config | [AppModelConfigResponse](#appmodelconfigresponse) | | No |
| name | string | | Yes |
| permission_keys | [ string ] | | No |
| site | [AppDetailSiteResponse](#appdetailsiteresponse) | | No |
@ -15519,6 +15543,35 @@ 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 |
@ -17147,7 +17200,7 @@ about. Stage 4 §4.2.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| 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 |
| children | [ { **"array_item"**: , **"children"**: [ object ], **"description"**: , **"file"**: , **"name"**: string, **"required"**: boolean, **"type"**: string, <br>**Available values:** "array", "boolean", "file", "number", "object", "string" } ] | | No |
| description | string | | No |
| type | [DeclaredOutputType](#declaredoutputtype) | | Yes |
@ -17176,7 +17229,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"**: , **"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 |
| children | [ { **"array_item"**: , **"children"**: [ object ], **"description"**: , **"file"**: , **"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 |
@ -21954,9 +22007,9 @@ The subscription constructor of the trigger provider
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| hash | string | | No |
| result | string | | No |
| updated_at | string | | No |
| hash | string | | Yes |
| result | string | | Yes |
| updated_at | integer | | Yes |
#### SystemConfigurationResponse
@ -21988,6 +22041,7 @@ 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 |
@ -22281,7 +22335,7 @@ Tool label
#### ToolParameter
Overrides type
Tool-specific parameter declaration and invocation-value normalization.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
@ -22294,6 +22348,7 @@ Overrides type
| 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 |

View File

@ -1582,6 +1582,7 @@ 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 |
@ -1733,7 +1734,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 | | Yes |
| icon_url | string | | No |
| input_placeholder | string | | No |
| privacy_policy | string | | No |
| prompt_public | boolean | | No |

View File

@ -1164,6 +1164,20 @@ 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
@ -1748,6 +1762,47 @@ 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,
@ -1767,6 +1822,26 @@ 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,

View File

@ -430,7 +430,11 @@ 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)
self._get_or_create_agent_app_debug_conversation(
agent=agent,
account_id=account_id,
draft_type=AgentConfigDraftType.DEBUG_BUILD,
)
return agent
def create_hidden_backing_app_for_workflow_agent(
@ -527,7 +531,11 @@ class AgentRosterService:
self._session.flush()
return backing_app.id
def _get_or_create_agent_app_debug_conversation(self, *, agent: Agent, account_id: str) -> str:
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."""
backing_app_id = self._ensure_workflow_agent_backing_app(agent=agent, account_id=account_id)
if not backing_app_id:
raise AgentNotFoundError()
@ -537,6 +545,7 @@ 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:
@ -570,6 +579,7 @@ class AgentRosterService:
agent_id=agent.id,
app_id=backing_app_id,
account_id=account_id,
draft_type=draft_type,
conversation_id=conversation_id,
)
)
@ -577,9 +587,15 @@ class AgentRosterService:
return conversation_id
def get_or_create_agent_app_debug_conversation_id(
self, *, tenant_id: str, agent_id: str, account_id: str, commit: bool = True
self,
*,
tenant_id: str,
agent_id: str,
account_id: str,
draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD,
commit: bool = True,
) -> str:
"""Return the current editor's debug conversation for an Agent App."""
"""Return the current editor's Build or Preview conversation for an Agent App."""
agent = self._session.scalar(
select(Agent).where(
@ -591,13 +607,24 @@ class AgentRosterService:
if agent is None:
raise AgentNotFoundError()
conversation_id = self._get_or_create_agent_app_debug_conversation(agent=agent, account_id=account_id)
conversation_id = self._get_or_create_agent_app_debug_conversation(
agent=agent,
account_id=account_id,
draft_type=draft_type,
)
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) -> str | None:
"""Return the current editor's existing debug conversation without creating or repairing rows."""
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."""
return self._session.scalar(
select(Conversation.id)
@ -606,6 +633,7 @@ 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,
@ -626,16 +654,21 @@ class AgentRosterService:
)
def refresh_agent_app_debug_conversation_id(
self, *, tenant_id: str, agent_id: str, account_id: str, commit: bool = True
self,
*,
tenant_id: str,
agent_id: str,
account_id: str,
draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD,
commit: bool = True,
) -> str:
"""Start a new console debug conversation for the current Agent App editor.
"""Start a new scoped console conversation for the current Agent App editor.
If this account already has a debug conversation mapping, the previous
If this account already has a mapping for the requested draft surface, 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 debug mapping is then repointed to the freshly created
conversation.
backend cleanup and then retired locally even when enqueueing fails. The
other draft surface is left untouched.
"""
agent = self._session.scalar(
@ -663,6 +696,7 @@ 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:
@ -672,6 +706,7 @@ class AgentRosterService:
agent_id=agent_id,
app_id=backing_app_id,
account_id=account_id,
draft_type=draft_type,
conversation_id=conversation_id,
)
)
@ -683,6 +718,7 @@ 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,
)
@ -699,6 +735,7 @@ class AgentRosterService:
tenant_id: str,
agent_id: str,
account_id: str,
draft_type: AgentConfigDraftType,
app_id: str,
conversation_id: str,
) -> None:
@ -727,7 +764,8 @@ 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}:{conversation_id}:debug-session-cleanup:"
f"{tenant_id}:{agent_id}:{account_id}:{draft_type.value}:{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'}"
@ -738,6 +776,7 @@ 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,
},
)
@ -772,9 +811,14 @@ class AgentRosterService:
)
def load_or_create_agent_app_debug_conversation_ids_by_agent_id(
self, *, tenant_id: str, agents: list[Agent], account_id: str
self,
*,
tenant_id: str,
agents: list[Agent],
account_id: str,
draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD,
) -> dict[str, str]:
"""Return per-account debug conversations for a page of Agent Apps."""
"""Return per-account scoped conversations for a page of Agent Apps."""
conversation_ids_by_agent_id: dict[str, str] = {}
changed = False
@ -784,6 +828,7 @@ 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:

View File

@ -5,6 +5,7 @@ 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,
@ -96,12 +97,14 @@ class BaseRequest:
logger.debug("Failed to generate traceparent header", exc_info=True)
with httpx.Client(mounts=mounts) as client:
# 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
# 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,
}
response = client.request(method, url, **request_kwargs)
@ -206,9 +209,8 @@ 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)

View File

@ -93,8 +93,20 @@ 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}"})
response = ssrf_proxy.post(
endpoint,
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
data=json.dumps(validation_payload),
)
except Exception as e:
raise ValueError(f"failed to connect to the endpoint: {endpoint}") from e
if response.status_code == 502:

View File

@ -186,6 +186,7 @@ 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:
@ -289,6 +290,7 @@ class FeatureService:
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]:

View File

@ -141,6 +141,29 @@ 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]

View File

@ -0,0 +1,502 @@
"""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}"),
)

View File

@ -1,33 +1,32 @@
"""Transport-only forwarding for the explicitly enabled KnowledgeFS Console operations.
"""Authorize and forward the explicitly enabled KnowledgeFS Console operations.
KnowledgeFS owns the request and response contract. This module binds short-lived
account and workspace identities, enforces Dify's coarse workspace policy, and
normalizes transport failures. Dify deliberately maintains a small product-facing
operation registry instead of exposing the full upstream OpenAPI surface. The
dedicated request path uses Dify's shared SSRF policy, never follows redirects,
bounds buffered responses, and rejects compressed responses.
The dedicated request path uses Dify's shared SSRF policy, never follows redirects,
bounds buffered responses, and rejects compressed streaming 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 Final, Literal, NamedTuple, Protocol
from typing import NamedTuple, Protocol
import httpx
import jwt
from configs import dify_config
from core.helper import ssrf_proxy
from core.rbac import RBACPermission, RBACResourceScope
from core.rbac import RBACResourceScope
from core.tools.errors import ToolSSRFError
from models import Account
from services.enterprise.rbac_service import RBACService
type KnowledgeFSMethod = Literal["DELETE", "GET", "PATCH", "POST", "PUT"]
type KnowledgeFSResponseKind = Literal["binary", "buffered", "stream"]
type KnowledgeFSRequiredScope = Literal["knowledge-spaces:read", "knowledge-spaces:write"]
from services.knowledge_fs_operations import (
KNOWLEDGE_FS_CONSOLE_OPERATIONS,
KnowledgeFSMethod,
KnowledgeFSOperation,
KnowledgeFSResponseKind,
)
_JWT_AUDIENCE = "knowledge-fs"
_JWT_ISSUER = "dify"
@ -35,56 +34,55 @@ _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]]: ...
@ -113,20 +111,29 @@ def authorize_knowledge_fs_request(
*,
account: Account,
tenant_id: str,
operation: KnowledgeFSOperation,
) -> None:
method: KnowledgeFSMethod,
path: str,
) -> KnowledgeFSAuthorization:
"""Enforce Dify's workspace policy before KFS performs resource authorization.
Args:
account: Authenticated Dify account with its current workspace role.
tenant_id: Current Dify workspace identifier.
operation: Dify-maintained KnowledgeFS operation and policy metadata.
method: Requested upstream HTTP method.
path: Requested relative KnowledgeFS path.
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.
"""
if operation.requires_dataset_editor and not account.is_dataset_editor:
raise KnowledgeFSAccessDeniedError("KnowledgeFS mutations require dataset edit access")
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 not RBACService.CheckAccess.check(
tenant_id,
account.id,
@ -134,6 +141,7 @@ 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(
@ -149,20 +157,62 @@ def proxy_knowledge_fs_request(
request_headers: _RequestHeaders | None = None,
) -> KnowledgeFSUpstreamResponse:
"""Authorize and forward one allowlisted KnowledgeFS request as a single use case."""
operation = get_knowledge_fs_operation(method, path)
authorize_knowledge_fs_request(
authorization = authorize_knowledge_fs_request(
account=account,
tenant_id=tenant_id,
operation=operation,
method=method,
path=path,
)
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=method,
path=path,
account_id=account_id,
method=operation.method,
path=operation.path,
tenant_id=tenant_id,
accept=accept,
content_type=content_type,

View File

@ -9,7 +9,9 @@ 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
@ -96,6 +98,39 @@ 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)

View File

@ -212,6 +212,17 @@ 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"]

View File

@ -1,15 +1,24 @@
"""Unit tests for the reset-encrypt-key-pair CLI command (#35396).
"""SQLite-backed 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 unittest.mock import MagicMock, patch
from types import SimpleNamespace
import pytest
from sqlalchemy import select
from sqlalchemy.orm import Session
import commands
from commands import system as system_commands
from models.provider import Provider, ProviderModel
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.tools import ApiToolProvider, BuiltinToolProvider, MCPToolProvider
@ -21,17 +30,60 @@ def _invoke_reset() -> int:
return 0
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
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 test_reset_aborts_when_not_self_hosted(monkeypatch, capsys):
@ -44,65 +96,73 @@ def test_reset_aborts_when_not_self_hosted(monkeypatch, capsys):
assert "only for SELF_HOSTED" in captured.out
def test_reset_purges_provider_and_tool_tables_for_each_tenant(monkeypatch, capsys):
@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:
"""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)
fake_tenant = MagicMock(id="tenant-abc", encrypt_public_key="old-key")
session = MagicMock()
session.scalars.return_value.all.return_value = [fake_tenant]
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_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()
exit_code = _invoke_reset()
captured = capsys.readouterr()
assert exit_code == 0
assert "tenant-abc" in captured.out
assert TENANT_ID in captured.out
# 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}"
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
)
def test_reset_iterates_all_tenants(monkeypatch, capsys):
@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:
"""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}")
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
_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()
fake_sessionmaker = MagicMock()
fake_sessionmaker.begin.return_value.__enter__.return_value = session
fake_sessionmaker.begin.return_value.__exit__.return_value = False
assert _invoke_reset() == 0
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}"
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() == []

View File

@ -37,6 +37,7 @@ 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
@ -148,32 +149,45 @@ def _configure_session_factory(_unit_test_engine):
configure_session_factory(_unit_test_engine, expire_on_commit=False)
def setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_owner):
"""
Helper to stub the tenant-owner execute result for service API app authentication.
def persist_service_api_tenant_owner(session: Session, tenant: Tenant, owner: Account) -> TenantAccountJoin:
"""Persist the owner identity resolved by 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().
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
The legacy name is retained temporarily for consumers on independent
conversion branches, but this helper no longer fabricates an execute result.
"""
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
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``."""
mock_db.session.execute.return_value.one_or_none.return_value = (mock_tenant, mock_owner)
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)
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,
)

View File

@ -53,6 +53,7 @@ from controllers.console.app.message import (
AgentMessageFeedbackApi,
AgentMessageSuggestedQuestionApi,
)
from models.agent import AgentConfigDraftType
from services.entities.agent_entities import ComposerSaveStrategy, ComposerVariant
@ -371,6 +372,7 @@ 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,
}
@ -544,8 +546,19 @@ def test_agent_app_copy_uses_agent_id_and_returns_agent_detail(
}
def test_agent_debug_conversation_refresh_uses_current_user(
app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str
@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,
) -> None:
agent_id = "00000000-0000-0000-0000-000000000001"
captured: dict[str, object] = {}
@ -557,7 +570,9 @@ def test_agent_debug_conversation_refresh_uses_current_user(
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"
"/console/api/agent/00000000-0000-0000-0000-000000000001/debug-conversation/refresh",
method="POST",
json=payload,
):
response = unwrap(AgentDebugConversationRefreshApi.post)(
AgentDebugConversationRefreshApi(), MagicMock(), "tenant-1", SimpleNamespace(id=account_id), agent_id
@ -567,7 +582,12 @@ def test_agent_debug_conversation_refresh_uses_current_user(
"debug_conversation_has_messages": False,
"debug_conversation_message_count": 0,
}
assert captured == {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id}
assert captured == {
"tenant_id": "tenant-1",
"agent_id": agent_id,
"account_id": account_id,
"draft_type": expected_draft_type,
}
def test_agent_publish_and_build_draft_routes_call_composer_service(
@ -1456,6 +1476,7 @@ 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
@ -1520,11 +1541,15 @@ 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",
lambda **kwargs: "debug-conversation-1",
resolve_debug_conversation,
)
monkeypatch.setattr(
completion_controller.helper, "compact_generate_response", lambda response: {"response": response}
@ -1544,6 +1569,7 @@ 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(
@ -1642,6 +1668,7 @@ 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]
@ -1649,13 +1676,26 @@ 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"}}
assert calls[1] == {
"get_or_create": {
"tenant_id": "tenant-1",
"agent_id": "agent-1",
"account_id": "account-1",
"draft_type": AgentConfigDraftType.DRAFT,
}
}
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"}
"get_or_create": {
"tenant_id": "tenant-1",
"agent_id": "backing-agent",
"account_id": "account-1",
"draft_type": AgentConfigDraftType.DEBUG_BUILD,
}
}

View File

@ -9,13 +9,14 @@ from unittest.mock import MagicMock
import pytest
from flask import Flask
from sqlalchemy import event
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
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 SystemFeatureModel, WebAppAuthModel
@ -66,6 +67,13 @@ def app() -> Iterator[Flask]:
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
@ -93,11 +101,34 @@ def transaction_events() -> TransactionEvents:
event.remove(Session, "after_rollback", record_rollback)
def _failed_result_after_starting_transaction(
service: app_import_module.AppDslService, *, app_id: str | None = None
) -> _Result:
service._session.begin()
return _Result(ImportStatus.FAILED, app_id=app_id)
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
class TestAppImportApi:
@ -110,15 +141,16 @@ class TestAppImportApi:
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
) -> None:
method = unwrap(api.post)
_install_features(monkeypatch, enabled=False)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda service, *_args, **_kwargs: _failed_result_after_starting_transaction(service, app_id=None),
app_id = _install_persisting_service_result(
monkeypatch,
method_name="import_app",
result=_Result(ImportStatus.FAILED, app_id=None),
)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
@ -126,6 +158,7 @@ class TestAppImportApi:
assert transaction_events.rollbacks == 1
assert transaction_events.commits == 0
_assert_app_persistence(sqlite_app_engine, app_id, persisted=False)
assert status == 400
assert response["status"] == ImportStatus.FAILED
@ -134,15 +167,16 @@ class TestAppImportApi:
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
) -> None:
method = unwrap(api.post)
_install_features(monkeypatch, enabled=False)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.PENDING),
app_id = _install_persisting_service_result(
monkeypatch,
method_name="import_app",
result=_Result(ImportStatus.PENDING),
)
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
@ -150,6 +184,7 @@ class TestAppImportApi:
assert transaction_events.commits == 1
assert transaction_events.rollbacks == 0
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
assert status == 202
assert response["status"] == ImportStatus.PENDING
@ -158,15 +193,16 @@ class TestAppImportApi:
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
) -> None:
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"),
app_id = _install_persisting_service_result(
monkeypatch,
method_name="import_app",
result=_Result(ImportStatus.COMPLETED, app_id="app-123"),
)
update_access = MagicMock()
monkeypatch.setattr(app_import_module.EnterpriseService.WebAppAuth, "update_app_access_mode", update_access)
@ -176,6 +212,7 @@ class TestAppImportApi:
assert transaction_events.commits == 1
assert transaction_events.rollbacks == 0
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
update_access.assert_called_once_with("app-123", "private")
assert status == 200
assert response["status"] == ImportStatus.COMPLETED
@ -185,6 +222,7 @@ class TestAppImportApi:
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
) -> None:
method = _unwrap(api.post)
@ -196,10 +234,10 @@ class TestAppImportApi:
lambda: (_make_account(), "tenant-1"),
)
monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.COMPLETED, app_id="app-123"),
app_id = _install_persisting_service_result(
monkeypatch,
method_name="import_app",
result=_Result(ImportStatus.COMPLETED, app_id="app-123"),
)
monkeypatch.setattr(
app_import_module,
@ -211,6 +249,7 @@ class TestAppImportApi:
response, status = method()
assert transaction_events.commits == 1
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
assert status == 200
assert response["permission_keys"] == ["app.acl.view_layout", "app.acl.edit"]
@ -219,6 +258,7 @@ class TestAppImportApi:
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
) -> None:
method = _unwrap(api.post)
@ -230,10 +270,10 @@ class TestAppImportApi:
lambda: (_make_account(), "tenant-1"),
)
monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True)
monkeypatch.setattr(
app_import_module.AppDslService,
"import_app",
lambda *_args, **_kwargs: _Result(ImportStatus.COMPLETED, app_id="app-123"),
app_id = _install_persisting_service_result(
monkeypatch,
method_name="import_app",
result=_Result(ImportStatus.COMPLETED, app_id="app-123"),
)
monkeypatch.setattr(
app_import_module,
@ -249,6 +289,7 @@ class TestAppImportApi:
response, status = method()
assert transaction_events.commits == 1
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
assert status == 200
assert response["permission_keys"] == []
@ -263,14 +304,15 @@ class TestAppImportConfirmApi:
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
) -> None:
method = unwrap(api.post)
monkeypatch.setattr(
app_import_module.AppDslService,
"confirm_import",
lambda service, *_args, **_kwargs: _failed_result_after_starting_transaction(service),
app_id = _install_persisting_service_result(
monkeypatch,
method_name="confirm_import",
result=_Result(ImportStatus.FAILED),
)
with app.test_request_context("/console/api/apps/imports/import-1/confirm", method="POST"):
@ -278,6 +320,7 @@ class TestAppImportConfirmApi:
assert transaction_events.rollbacks == 1
assert transaction_events.commits == 0
_assert_app_persistence(sqlite_app_engine, app_id, persisted=False)
assert status == 400
assert response["status"] == ImportStatus.FAILED
@ -286,6 +329,7 @@ class TestAppImportConfirmApi:
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
) -> None:
method = _unwrap(api.post)
@ -304,10 +348,10 @@ class TestAppImportConfirmApi:
),
)
monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True)
monkeypatch.setattr(
app_import_module.AppDslService,
"confirm_import",
lambda *_args, **_kwargs: _Result(ImportStatus.COMPLETED, app_id="app-456"),
app_id = _install_persisting_service_result(
monkeypatch,
method_name="confirm_import",
result=_Result(ImportStatus.COMPLETED, app_id="app-456"),
)
monkeypatch.setattr(
app_import_module,
@ -319,6 +363,7 @@ class TestAppImportConfirmApi:
response, status = method(import_id="import-1")
assert transaction_events.commits == 1
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
assert status == 200
assert response["permission_keys"] == ["app.acl.view_layout", "app.acl.edit"]
@ -327,6 +372,7 @@ class TestAppImportConfirmApi:
api,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_app_engine: Engine,
transaction_events: TransactionEvents,
) -> None:
method = _unwrap(api.post)
@ -345,10 +391,10 @@ class TestAppImportConfirmApi:
),
)
monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True)
monkeypatch.setattr(
app_import_module.AppDslService,
"confirm_import",
lambda *_args, **_kwargs: _Result(ImportStatus.COMPLETED, app_id="app-456"),
app_id = _install_persisting_service_result(
monkeypatch,
method_name="confirm_import",
result=_Result(ImportStatus.COMPLETED, app_id="app-456"),
)
monkeypatch.setattr(
app_import_module,
@ -360,6 +406,7 @@ class TestAppImportConfirmApi:
response, status = method(import_id="import-1")
assert transaction_events.commits == 1
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
assert status == 200
assert response["permission_keys"] == []

View File

@ -1,5 +1,6 @@
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
@ -102,12 +103,14 @@ def test_structured_output_generate_exceptions(app: Flask, monkeypatch: pytest.M
method(api, "t1")
def test_instruction_generate_exceptions(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_instruction_generate_exceptions(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
) -> None:
api = generator_module.InstructionGenerateApi()
method = unwrap(api.post)
from types import SimpleNamespace
session = SimpleNamespace()
exceptions_to_test = [
(ProviderTokenNotInitError("token error"), generator_module.ProviderNotInitializeError),
@ -135,4 +138,4 @@ def test_instruction_generate_exceptions(app: Flask, monkeypatch: pytest.MonkeyP
},
):
with pytest.raises(expected_exception):
method(api, session, "t1")
method(api, sqlite_session, "t1")

View File

@ -11,7 +11,7 @@ import pytest
from flask import Flask
from sqlalchemy import func, select
from sqlalchemy.engine import Engine
from sqlalchemy.orm import object_session, sessionmaker
from sqlalchemy.orm import Session, object_session, sessionmaker
from controllers.common import session as controller_session
from controllers.console.app import model_config as model_config_module
@ -29,11 +29,14 @@ 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)
@ -45,14 +48,16 @@ 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"]
@ -62,18 +67,16 @@ 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, session, "t1", "u1", app_model=app_model)
response = method(api, sqlite_session, "t1", "u1", app_model=app_model)
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 send_mock.call_args.kwargs["session"] is sqlite_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"
@ -160,7 +163,11 @@ def test_post_uses_one_session_and_rolls_back_when_signal_fails(
assert config_count == 1
def test_post_encrypts_agent_tool_parameters(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
@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."""
api = model_config_module.ModelConfigResource()
method = unwrap(api.post)
@ -174,6 +181,7 @@ def test_post_encrypts_agent_tool_parameters(app: Flask, monkeypatch: pytest.Mon
_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,
@ -190,8 +198,8 @@ def test_post_encrypts_agent_tool_parameters(app: Flask, monkeypatch: pytest.Mon
}
)
session = MagicMock()
session.scalar.return_value = original_config
sqlite_session.add(original_config)
sqlite_session.commit()
monkeypatch.setattr(
model_config_module.AppModelConfigService,
@ -236,11 +244,11 @@ def test_post_encrypts_agent_tool_parameters(app: Flask, monkeypatch: pytest.Mon
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, session, "t1", "u1", app_model=app_model)
response = method(api, sqlite_session, "t1", "u1", app_model=app_model)
stored_config = session.add.call_args[0][0]
stored_config = sqlite_session.get(AppModelConfig, app_model.app_model_config_id)
assert stored_config is not None
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"

View File

@ -14,6 +14,7 @@ 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
@ -202,6 +203,50 @@ 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")

View File

@ -1,4 +1,4 @@
"""Testcontainers integration tests for OAuth controller endpoints."""
"""Unit tests for OAuth controller endpoints."""
from __future__ import annotations
@ -16,15 +16,10 @@ 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"),
[
@ -65,10 +60,6 @@ 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()
@ -181,10 +172,6 @@ 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"""
@ -263,10 +250,12 @@ 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,
@ -280,13 +269,20 @@ 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_by_token.return_value = {"email": "user@example.com"}
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"]
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_by_token.assert_called_once_with(token="invite123")
mock_register_service.get_invitation_if_token_valid.assert_called_once_with(
None, None, "invite123", session=ANY
)
mock_redirect.assert_called_once_with("http://localhost:3000/signin/invite-settings?invite_token=invite123")
@pytest.mark.parametrize(
@ -448,10 +444,6 @@ 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")
@ -468,39 +460,25 @@ class TestAccountGeneration:
self,
mock_account_model,
mock_get_account,
flask_req_ctx_with_containers,
app: Flask,
user_info: OAuthUserInfo,
mock_account,
):
# 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()
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 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()
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
result = _get_account_by_openid_or_email("github", user_info)
assert result == mock_account
mock_get_account.assert_called_once()
@pytest.mark.parametrize(
("allow_register", "existing_account", "should_create"),

View File

@ -1,5 +1,5 @@
import urllib.parse
from unittest.mock import MagicMock, patch
from unittest.mock import ANY, MagicMock, patch
import pytest
from flask import Flask
@ -91,3 +91,102 @@ 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()

View File

@ -1,12 +1,14 @@
"""Testcontainers integration tests for password reset authentication flows."""
"""Unit tests for password reset controller 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
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from controllers.console.auth.error import (
EmailCodeError,
@ -21,47 +23,58 @@ from controllers.console.auth.forgot_password import (
ForgotPasswordSendEmailApi,
)
from controllers.console.error import AccountNotFound, EmailSendIpLimitError
from tests.test_containers_integration_tests.controllers.console.helpers import ensure_dify_setup
from models.account import Account, Tenant, TenantAccountJoin
from services.feature_service import 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(enable_email_password_login=True),
)
class TestForgotPasswordSendEmailApi:
"""Test cases for sending password reset emails."""
@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
@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_success(
self,
mock_get_features,
mock_send_email,
mock_get_account,
mock_is_ip_limit,
app: Flask,
mock_account,
sqlite_session: Session,
):
# 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 app.test_request_context(
"/forgot-password", method="POST", json={"email": "test@example.com", "language": "en-US"}
with (
_bind_database_session(sqlite_session),
app.test_request_context(
"/forgot-password", method="POST", json={"email": "test@example.com", "language": "en-US"}
),
):
api = ForgotPasswordSendEmailApi()
response = api.post()
@ -98,20 +111,17 @@ 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.
@ -122,13 +132,14 @@ 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 app.test_request_context(
"/forgot-password", method="POST", json={"email": "test@example.com", "language": language_input}
with (
_bind_database_session(sqlite_session),
app.test_request_context(
"/forgot-password", method="POST", json={"email": "test@example.com", "language": language_input}
),
):
api = ForgotPasswordSendEmailApi()
api.post()
@ -141,11 +152,6 @@ 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")
@ -153,10 +159,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,
mock_generate_token,
mock_revoke_token,
mock_get_data,
mock_reset_rate_limit: MagicMock,
mock_generate_token: MagicMock,
mock_revoke_token: MagicMock,
mock_get_data: MagicMock,
mock_is_rate_limit,
app: Flask,
):
@ -200,10 +206,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,
mock_generate_token,
mock_revoke_token,
mock_get_data,
mock_reset_rate_limit: MagicMock,
mock_generate_token: MagicMock,
mock_revoke_token: MagicMock,
mock_get_data: MagicMock,
mock_is_rate_limit,
app: Flask,
):
@ -325,33 +331,15 @@ class TestForgotPasswordCheckApi:
class TestForgotPasswordResetApi:
"""Test cases for resetting password with verified token."""
@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
@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")
@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_get_tenants,
mock_db,
mock_get_account,
mock_revoke_token,
mock_get_data,
mock_revoke_token: MagicMock,
mock_get_data: MagicMock,
app: Flask,
mock_account,
sqlite_session: Session,
):
"""
Test successful password reset.
@ -363,25 +351,39 @@ 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 app.test_request_context(
"/forgot-password/resets",
method="POST",
json={"token": "valid_token", "new_password": "NewPass123!", "password_confirm": "NewPass123!"},
):
api = ForgotPasswordResetApi()
response = api.post()
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)
# 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
@patch("controllers.console.auth.forgot_password.AccountService.get_reset_password_data")
def test_reset_password_mismatch(self, mock_get_data, app: Flask):
def test_reset_password_mismatch(self, app: Flask):
"""
Test password reset with mismatched passwords.
@ -389,9 +391,6 @@ 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",
@ -445,10 +444,12 @@ 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")
@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):
def test_reset_password_account_not_found(
self, mock_revoke_token, mock_get_data, app: Flask, sqlite_session: Session
):
"""
Test password reset for non-existent account.
@ -457,13 +458,15 @@ class TestForgotPasswordResetApi:
"""
# Arrange
mock_get_data.return_value = {"email": "nonexistent@example.com", "phase": "reset"}
mock_get_account.return_value = None
# Act & Assert
with app.test_request_context(
"/forgot-password/resets",
method="POST",
json={"token": "token", "new_password": "NewPass123!", "password_confirm": "NewPass123!"},
with (
_bind_database_session(sqlite_session),
app.test_request_context(
"/forgot-password/resets",
method="POST",
json={"token": "token", "new_password": "NewPass123!", "password_confirm": "NewPass123!"},
),
):
api = ForgotPasswordResetApi()
with pytest.raises(AccountNotFound):

View File

@ -1,3 +1,9 @@
"""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
@ -7,6 +13,8 @@ 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
@ -73,90 +81,80 @@ 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
app, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine
) -> None:
api = module.PublishedAllRagPipelineApi()
handler = unwrap_all(api.get)
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()
session_state: dict[str, Session] = {}
base_workflow = _make_workflow()
class _Workflow:
def __getattr__(self, name: str):
assert session_state["open"] is True
assert session_state["session"].in_transaction() is True
return getattr(base_workflow, name)
monkeypatch.setattr(module, "db", SimpleNamespace(engine=object(), session=lambda: object()))
monkeypatch.setattr(module, "sessionmaker", lambda *_args, **_kwargs: _SessionMaker())
def _get_all_published_workflow(**kwargs):
session_state["session"] = kwargs["session"]
return [_Workflow()], False
monkeypatch.setattr(
module,
"RagPipelineService",
lambda *_args, **_kwargs: SimpleNamespace(get_all_published_workflow=lambda **_kwargs: ([_Workflow()], False)),
lambda *_args, **_kwargs: SimpleNamespace(get_all_published_workflow=_get_all_published_workflow),
)
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 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())
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) -> None:
def test_rag_pipeline_workflow_patch_serializes_response_model(
app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine
) -> None:
workflow = _make_workflow(marked_name="Updated release")
captured_session: dict[str, Session] = {}
class _SessionContext:
def __enter__(self):
return object()
def _update_workflow(**kwargs):
captured_session["session"] = kwargs["session"]
assert kwargs["session"].in_transaction() is True
return workflow
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=lambda **_kwargs: workflow),
lambda *_args, **_kwargs: SimpleNamespace(update_workflow=_update_workflow),
)
payload: dict[str, object] = {"marked_name": "Updated release"}
api = module.RagPipelineByIdApi()
handler = unwrap_all(api.patch)
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 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",
)
assert captured_session["session"].in_transaction() is False
assert response["id"] == "workflow-1"
assert response["marked_name"] == "Updated release"
assert response["hash"] == "hash-1"

View File

@ -8,6 +8,8 @@ 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
@ -36,7 +38,9 @@ def _snippet(**overrides) -> CustomizedSnippet:
@pytest.fixture(autouse=True)
def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch) -> None:
def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None:
snippet_session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
def factory():
try:
return snippet_workflow_module.SnippetService(snippet_workflow_module._snippet_session_maker())
@ -44,7 +48,7 @@ def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch) -> None:
return snippet_workflow_module.SnippetService()
monkeypatch.setattr(snippet_workflow_module, "_snippet_service", factory)
monkeypatch.setattr(snippet_workflow_module, "_snippet_session_maker", Mock(return_value=Mock()))
monkeypatch.setattr(snippet_workflow_module, "_snippet_session_maker", lambda: snippet_session_maker)
def test_get_snippet_requires_snippet_id(app):
@ -150,28 +154,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
def test_published_workflow_post_returns_400_when_publish_fails(app: Flask, monkeypatch: pytest.MonkeyPatch) -> 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:
user = _account("account-1")
snippet = _snippet()
merged_snippet = _snippet()
session = SimpleNamespace(merge=Mock(return_value=merged_snippet), commit=Mock())
sqlite_session.add(snippet)
sqlite_session.commit()
class SessionContext:
def __init__(self, engine):
self.engine = engine
def fail_publish(*, session: Session, snippet: CustomizedSnippet, account: Account):
snippet.name = "Uncommitted name"
session.add(snippet)
raise ValueError("No valid workflow found.")
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, "db", SimpleNamespace(engine=sqlite_engine))
monkeypatch.setattr(
snippet_workflow_module,
"SnippetService",
lambda: SimpleNamespace(publish_workflow=Mock(side_effect=ValueError("No valid workflow found."))),
lambda: SimpleNamespace(publish_workflow=Mock(side_effect=fail_publish)),
)
api = snippet_workflow_module.SnippetPublishedWorkflowApi()
@ -182,7 +186,8 @@ def test_published_workflow_post_returns_400_when_publish_fails(app: Flask, monk
assert status_code == 400
assert response == {"message": "No valid workflow found."}
session.commit.assert_not_called()
sqlite_session.refresh(snippet)
assert snippet.name == "Snippet"
def test_default_block_configs_delegates_to_service(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
@ -203,7 +208,11 @@ 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) -> None:
def test_list_published_snippet_workflows_includes_input_fields(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
) -> None:
workflow = SimpleNamespace(
id="workflow-1",
graph_dict={"nodes": [], "edges": []},
@ -224,18 +233,7 @@ def test_list_published_snippet_workflows_includes_input_fields(app: Flask, monk
input_fields = [{"variable": "query", "type": "text"}]
snippet = _snippet(input_fields=json.dumps(input_fields))
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, "db", SimpleNamespace(engine=sqlite_engine))
monkeypatch.setattr(
snippet_workflow_module,
"SnippetService",
@ -364,8 +362,11 @@ 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
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
) -> None:
workflow = SimpleNamespace(
id="workflow-1",
@ -387,21 +388,15 @@ 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))
session = SimpleNamespace()
update_workflow = Mock(return_value=workflow)
sqlite_session.add(snippet)
sqlite_session.commit()
class TransactionContext:
def __enter__(self):
return session
def update_persisted_snippet(*, session: Session, snippet: CustomizedSnippet, **_kwargs):
merged_snippet = session.merge(snippet)
merged_snippet.description = "Updated in transaction"
return workflow
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()))
update_workflow = Mock(side_effect=update_persisted_snippet)
monkeypatch.setattr(
snippet_workflow_module,
"SnippetService",
@ -418,16 +413,18 @@ def test_update_published_snippet_workflow_returns_updated_workflow(
):
response = handler(api, user, snippet, workflow_id="workflow-1")
update_workflow.assert_called_once_with(
session=session,
snippet=snippet,
workflow_id="workflow-1",
account=user,
data={"marked_name": "v1", "marked_comment": "first version"},
)
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"}
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:
@ -441,26 +438,25 @@ def test_update_published_snippet_workflow_returns_400_when_no_fields(app: Flask
assert response == {"message": "No valid fields to update"}
def test_update_published_snippet_workflow_raises_not_found(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
@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:
user = _account("account-1")
snippet = _snippet()
sqlite_session.add(snippet)
sqlite_session.commit()
class TransactionContext:
def __enter__(self):
return SimpleNamespace()
def update_missing_workflow(*, session: Session, snippet: CustomizedSnippet, **_kwargs):
merged_snippet = session.merge(snippet)
merged_snippet.name = "Rolled back name"
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(return_value=None)),
lambda: SimpleNamespace(update_workflow=Mock(side_effect=update_missing_workflow)),
)
api = snippet_workflow_module.SnippetWorkflowByIdApi()
@ -474,6 +470,9 @@ def test_update_published_snippet_workflow_raises_not_found(app: Flask, monkeypa
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()

View File

@ -1,15 +1,29 @@
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 core.workflow.variable_prefixes import CONVERSATION_VARIABLE_NODE_ID, SYSTEM_VARIABLE_NODE_ID
from graphon.variables import StringSegment
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(
@ -21,8 +35,31 @@ 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):
def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch) -> None:
def factory():
service_factory = module.SnippetService
if isinstance(service_factory, type):
@ -33,33 +70,69 @@ def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch):
@pytest.fixture
def app():
def app() -> Flask:
app = Flask("test_snippet_workflow_draft_variable")
app.config["TESTING"] = True
return app
def test_ensure_snippet_draft_variable_row_allowed_rejects_system_variable():
variable = SimpleNamespace(node_id=SYSTEM_VARIABLE_NODE_ID)
@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,
)
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():
variable = SimpleNamespace(node_id=CONVERSATION_VARIABLE_NODE_ID)
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"),
)
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():
variable = SimpleNamespace(node_id="llm-1")
def test_ensure_snippet_draft_variable_row_allowed_accepts_canvas_node_variable() -> None:
variable = _make_node_variable("var-1")
module._ensure_snippet_draft_variable_row_allowed(variable=variable, variable_id="var-1")
def test_conversation_variables_returns_empty_list(app: Flask):
def test_conversation_variables_returns_empty_list(app: Flask) -> None:
api = module.SnippetConversationVariableCollectionApi()
handler = unwrap(api.get)
@ -69,7 +142,7 @@ def test_conversation_variables_returns_empty_list(app: Flask):
assert result == WorkflowDraftVariableList(variables=[])
def test_system_variables_returns_empty_list(app: Flask):
def test_system_variables_returns_empty_list(app: Flask) -> None:
api = module.SnippetSystemVariableCollectionApi()
handler = unwrap(api.get)
@ -79,12 +152,17 @@ def test_system_variables_returns_empty_list(app: Flask):
assert result == WorkflowDraftVariableList(variables=[])
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)
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)
api = module.SnippetWorkflowVariableCollectionApi()
handler = unwrap(api.delete)
@ -92,11 +170,14 @@ def test_delete_variable_collection_deletes_current_user_variables(app: Flask, m
response = handler(api, _make_account(), snippet=SimpleNamespace(id="snippet-1"))
assert response.status_code == 204
draft_var_service.delete_user_workflow_variables.assert_called_once_with("snippet-1", user_id="user-1")
db_session.commit.assert_called_once()
assert _variable_ids(sqlite_engine) == {other_user.id, other_snippet.id}
assert not controller_sessions().in_transaction()
def test_variable_collection_get_raises_when_draft_workflow_missing(app: Flask, monkeypatch: pytest.MonkeyPatch):
def test_variable_collection_get_raises_when_draft_workflow_missing(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(
module,
"SnippetService",
@ -111,47 +192,37 @@ def test_variable_collection_get_raises_when_draft_workflow_missing(app: Flask,
handler(api, _make_account(), snippet=SimpleNamespace(id="snippet-1"))
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)),
)
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)
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 result is variables
list_node_variables.assert_called_once_with("snippet-1", "llm-1", user_id="user-1")
assert [variable.id for variable in result.variables] == [matching.id]
assert controller_sessions().get_bind() is not None
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)
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)
api = module.SnippetNodeVariableCollectionApi()
handler = unwrap(api.delete)
@ -159,83 +230,102 @@ def test_node_variable_collection_delete_deletes_node_variables(app: Flask, monk
response = handler(api, _make_account(), snippet=SimpleNamespace(id="snippet-1"), node_id="llm-1")
assert response.status_code == 204
delete_node_variables.assert_called_once_with("snippet-1", "llm-1", user_id="user-1")
db_session.commit.assert_called_once()
assert _variable_ids(sqlite_engine) == {other_node.id, other_user.id}
assert not controller_sessions().in_transaction()
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 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 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)
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()
assert result.id == variable.id
assert result.app_id == "snippet-1"
assert commits == []
assert session.in_transaction()
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))
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)
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="var-1")
response = handler(
api,
_make_account(),
snippet=SimpleNamespace(id="snippet-1"),
variable_id=variable.id,
)
assert response.status_code == 204
delete_variable.assert_called_once_with(variable)
db_session.commit.assert_called_once()
assert _variable_ids(sqlite_engine) == {retained.id}
assert not controller_sessions().in_transaction()
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))
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)
monkeypatch.setattr(
module,
"SnippetService",
Mock(return_value=SimpleNamespace(get_draft_workflow=Mock(return_value=draft_workflow))),
Mock(return_value=SimpleNamespace(get_draft_workflow=Mock(return_value=SimpleNamespace(id="workflow-1")))),
)
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="var-1")
response = handler(
api,
_make_account(),
snippet=SimpleNamespace(id="snippet-1"),
variable_id=variable.id,
)
assert response.status_code == 204
draft_var_service.reset_variable.assert_called_once_with(draft_workflow, variable)
db_session.commit.assert_called_once()
assert _variable_ids(sqlite_engine) == set()
assert not controller_sessions().in_transaction()
def test_environment_variables_returns_workflow_environment_variables(app: Flask, monkeypatch: pytest.MonkeyPatch):
def test_environment_variables_returns_workflow_environment_variables(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
) -> None:
env_var = SimpleNamespace(
id="env-1",
name="API_KEY",

View File

@ -1,9 +1,11 @@
from collections.abc import Iterator
from types import SimpleNamespace
from unittest.mock import MagicMock, PropertyMock, patch
from unittest.mock import PropertyMock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from sqlalchemy import Engine
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from werkzeug.exceptions import Forbidden
import controllers.console.tag.tags as module
@ -16,15 +18,12 @@ 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.
@ -41,6 +40,26 @@ 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(
@ -66,11 +85,16 @@ def readonly_user():
@pytest.fixture
def tag():
tag = MagicMock()
def tag(sqlite_db_session: scoped_session[Session]):
tag = Tag(
tenant_id="tenant-1",
name="test-tag",
type=TagType.KNOWLEDGE,
created_by="user-1",
)
tag.id = "tag-1"
tag.name = "test-tag"
tag.type = TagType.KNOWLEDGE
sqlite_db_session.add(tag)
sqlite_db_session.commit()
return tag
@ -111,7 +135,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):
def test_get_snippet_tags(self, app: Flask, sqlite_engine: Engine):
api = TagListApi()
method = unwrap(api.get)
@ -131,7 +155,9 @@ class TestTagListApi:
):
result, status = method(api, "tenant-1")
get_tags_mock.assert_called_once_with("snippet", "tenant-1", None, session=SessionMatcher())
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)
assert status == 200
assert result == [{"id": "1", "name": "snippet-tag", "type": "snippet", "binding_count": "1"}]
@ -200,7 +226,7 @@ class TestTagListApi:
class TestTagUpdateDeleteApi:
def test_patch_success(self, app: Flask, admin_user, tag, payload_patch):
def test_patch_success(self, app: Flask, admin_user, tag, payload_patch, sqlite_engine: Engine):
api = TagUpdateDeleteApi()
method = unwrap(api.patch)
@ -224,7 +250,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 session == SessionMatcher()
_assert_sqlite_session(session, sqlite_engine)
assert result["binding_count"] == "3"
def test_patch_forbidden(self, app: Flask, readonly_user, payload_patch):
@ -240,7 +266,7 @@ class TestTagUpdateDeleteApi:
with pytest.raises(Forbidden):
method(api, readonly_user, "tag-1")
def test_delete_success(self, app: Flask, admin_user):
def test_delete_success(self, app: Flask, admin_user, sqlite_engine: Engine):
api = TagUpdateDeleteApi()
method = unwrap(api.delete)
@ -250,12 +276,30 @@ class TestTagUpdateDeleteApi:
):
result, status = method(api, "tag-1")
delete_mock.assert_called_once_with("tag-1", SessionMatcher())
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 status == 204
def test_delete_snippet_tag_checks_type_in_current_tenant(self, app: Flask, admin_user):
def test_delete_snippet_tag_checks_type_in_current_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="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("/"),
@ -264,13 +308,11 @@ 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",
@ -278,7 +320,49 @@ class TestTagUpdateDeleteApi:
scene=module.RBACPermission.SNIPPETS_CREATE_AND_MODIFY,
resource_required=False,
)
delete_mock.assert_called_once_with("tag-1", SessionMatcher())
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)
assert result == ""
assert status == 204

View File

@ -1,27 +1,16 @@
"""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
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)
from models.model import DifySetup
def test_get_init_status_finished(monkeypatch: pytest.MonkeyPatch) -> None:
@ -85,11 +74,15 @@ def test_get_init_validate_status_validated_session(app: Flask, monkeypatch: pyt
assert init_validate.get_init_validate_status() is True
def test_get_init_validate_status_setup_exists(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
@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:
monkeypatch.setattr(init_validate.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setenv("INIT_PASSWORD", "expected")
monkeypatch.setattr(init_validate, "Session", lambda *_args, **_kwargs: _SessionStub(True))
monkeypatch.setattr(init_validate, "db", SimpleNamespace(engine=object()))
monkeypatch.setattr(init_validate, "db", SimpleNamespace(engine=sqlite_session.get_bind()))
sqlite_session.add(DifySetup(version="test-version"))
sqlite_session.commit()
app.secret_key = "test-secret"
with app.test_request_context("/console/api/init", method="GET"):
@ -97,11 +90,13 @@ def test_get_init_validate_status_setup_exists(app: Flask, monkeypatch: pytest.M
assert init_validate.get_init_validate_status() is True
def test_get_init_validate_status_not_validated(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
@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:
monkeypatch.setattr(init_validate.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setenv("INIT_PASSWORD", "expected")
monkeypatch.setattr(init_validate, "Session", lambda *_args, **_kwargs: _SessionStub(False))
monkeypatch.setattr(init_validate, "db", SimpleNamespace(engine=object()))
monkeypatch.setattr(init_validate, "db", SimpleNamespace(engine=sqlite_session.get_bind()))
app.secret_key = "test-secret"
with app.test_request_context("/console/api/init", method="GET"):

View File

@ -7,9 +7,11 @@ 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,
@ -22,15 +24,18 @@ 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_proxy import (
KnowledgeFSAccessDeniedError,
KnowledgeFSConfigurationError,
from services.knowledge_fs_operations import (
KnowledgeFSMethod,
KnowledgeFSOperation,
KnowledgeFSResponseKind,
)
from services.knowledge_fs_proxy import (
KnowledgeFSAccessDeniedError,
KnowledgeFSConfigurationError,
KnowledgeFSRouteNotAllowedError,
KnowledgeFSUpstreamResponse,
get_knowledge_fs_operation,
@ -46,6 +51,7 @@ 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(
@ -55,7 +61,7 @@ def _upstream(
response_kind=kind,
required_scope="knowledge-spaces:read",
rbac_permission=RBACPermission.DATASET_READONLY,
requires_dataset_editor=False,
legacy_role="reader",
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=(),
@ -66,6 +72,7 @@ def _upstream(
"x-session-id",
),
response_media_types=(),
error_status_map=error_status_map,
)
return KnowledgeFSUpstreamResponse(response, kind, operation)
@ -92,7 +99,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",
unwrap(_proxy_knowledge_fs_non_get),
lambda method, path: _proxy_request(method, path),
)
@ -118,10 +125,61 @@ 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_get")
assert options_endpoint.endswith("proxy_knowledge_fs_options")
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)
@ -235,10 +293,8 @@ 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)
monkeypatch.setattr(
"services.knowledge_fs_proxy.RBACService.CheckAccess.check",
MagicMock(return_value=True),
)
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=True, limit=10)),
@ -248,14 +304,19 @@ 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_request", proxy)
monkeypatch.setattr("controllers.console.knowledge_fs_proxy._proxy_authorized_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_with("POST", "knowledge-spaces")
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()
def test_denied_write_does_not_consume_the_workspace_rate_limit(
@ -395,6 +456,65 @@ 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,
@ -565,6 +685,43 @@ 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,
@ -619,6 +776,7 @@ 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):

View File

@ -2,19 +2,75 @@
Unit tests for inner_api plugin decorators
"""
from collections.abc import Iterator
from types import SimpleNamespace
from typing import Any
from unittest.mock import MagicMock, patch
import pytest
from flask import Flask
from pydantic import ValidationError
from sqlalchemy import Engine, event, select
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from controllers.inner_api.plugin import wraps as wraps_module
from controllers.inner_api.plugin.wraps import (
TenantUserPayload,
get_user,
get_user_tenant,
plugin_data,
)
from models.account import Tenant
from models.base import TypeBase
from models.enums import EndUserType
from models.model import DefaultEndUserSessionID, EndUser
@pytest.fixture
def sqlite_plugin_engine(
sqlite_engine: Engine,
monkeypatch: pytest.MonkeyPatch,
) -> Iterator[Engine]:
tables = [TypeBase.metadata.tables[model.__tablename__] for model in (Tenant, EndUser)]
TypeBase.metadata.create_all(sqlite_engine, tables=tables)
session_registry = scoped_session(sessionmaker(bind=sqlite_engine, expire_on_commit=False))
monkeypatch.setattr(
wraps_module,
"db",
SimpleNamespace(engine=sqlite_engine, session=session_registry),
)
try:
yield sqlite_engine
finally:
session_registry.remove()
def _persist_tenant(sqlite_engine: Engine, *, tenant_id: str = "tenant123") -> Tenant:
tenant = Tenant(name=f"Tenant {tenant_id}")
tenant.id = tenant_id
with Session(sqlite_engine) as session, session.begin():
session.add(tenant)
return tenant
def _persist_end_user(
sqlite_engine: Engine,
*,
tenant_id: str = "tenant123",
user_id: str,
session_id: str,
is_anonymous: bool = False,
) -> EndUser:
user = EndUser(
id=user_id,
tenant_id=tenant_id,
type=EndUserType.SERVICE_API,
is_anonymous=is_anonymous,
session_id=session_id,
)
with Session(sqlite_engine) as session, session.begin():
session.add(user)
return user
class TestTenantUserPayload:
@ -41,185 +97,143 @@ class TestTenantUserPayload:
class TestGetUser:
"""Test get_user function"""
@patch("controllers.inner_api.plugin.wraps.select")
@patch("controllers.inner_api.plugin.wraps.EndUser")
@patch("controllers.inner_api.plugin.wraps.sessionmaker")
@patch("controllers.inner_api.plugin.wraps.db")
def test_should_return_existing_user_by_id(
self, mock_db, mock_sessionmaker, mock_enduser_class, mock_select, app: Flask
):
def test_should_return_existing_user_by_id(self, sqlite_plugin_engine: Engine, app: Flask):
"""Test returning existing user when found by ID"""
# Arrange
mock_user = MagicMock()
mock_user.id = "user123"
mock_session = MagicMock()
mock_sessionmaker.return_value.begin.return_value.__enter__.return_value = mock_session
mock_session.scalar.return_value = mock_user
mock_query = MagicMock()
mock_select.return_value.where.return_value.limit.return_value = mock_query
_persist_end_user(
sqlite_plugin_engine,
user_id="user123",
session_id="existing-session",
)
# Act
with app.app_context():
result = get_user("tenant123", "user123")
# Assert
assert result == mock_user
mock_session.scalar.assert_called_once()
assert result.id == "user123"
assert result.tenant_id == "tenant123"
@patch("controllers.inner_api.plugin.wraps.select")
@patch("controllers.inner_api.plugin.wraps.EndUser")
@patch("controllers.inner_api.plugin.wraps.sessionmaker")
@patch("controllers.inner_api.plugin.wraps.db")
def test_should_not_resolve_non_anonymous_users_across_tenants(
self,
mock_db,
mock_sessionmaker,
mock_enduser_class,
mock_select,
sqlite_plugin_engine: Engine,
app: Flask,
):
"""Test that explicit user IDs remain scoped to the current tenant."""
# Arrange
mock_session = MagicMock()
mock_sessionmaker.return_value.begin.return_value.__enter__.return_value = mock_session
mock_session.scalar.return_value = None
mock_new_user = MagicMock()
mock_new_user.tenant_id = "tenant-current"
mock_enduser_class.return_value = mock_new_user
_persist_end_user(
sqlite_plugin_engine,
tenant_id="tenant-foreign",
user_id="foreign-user-id",
session_id="foreign-session",
)
# Act
with app.app_context():
result = get_user("tenant-current", "foreign-user-id")
# Assert
assert result == mock_new_user
mock_session.get.assert_not_called()
# Non-anonymous miss now tries id, then session_id fallback (see
# #36736); both miss in this tenant → fall through to create.
assert mock_session.scalar.call_count == 2
mock_session.add.assert_called_once_with(mock_new_user)
assert result.id != "foreign-user-id"
assert result.tenant_id == "tenant-current"
assert result.session_id == "foreign-user-id"
with Session(sqlite_plugin_engine) as session:
current_tenant_users = session.scalars(select(EndUser).where(EndUser.tenant_id == "tenant-current")).all()
assert [user.id for user in current_tenant_users] == [result.id]
@patch("controllers.inner_api.plugin.wraps.select")
@patch("controllers.inner_api.plugin.wraps.EndUser")
@patch("controllers.inner_api.plugin.wraps.sessionmaker")
@patch("controllers.inner_api.plugin.wraps.db")
def test_should_return_existing_user_by_session_id_fallback_for_non_anonymous(
self, mock_db, mock_sessionmaker, mock_enduser_class, mock_select, app: Flask
self,
sqlite_plugin_engine: Engine,
app: Flask,
):
"""Non-anonymous user_id misses on EndUser.id but hits on
EndUser.session_id this is the plugin-daemon Reverse Invocation
case where the daemon sends a stable session-derived UUID that
was written into session_id on the first call. See #36736.
"""
# Arrange
mock_user = MagicMock()
mock_session = MagicMock()
mock_sessionmaker.return_value.begin.return_value.__enter__.return_value = mock_session
# First scalar (id lookup) returns None, second (session_id fallback) hits.
mock_session.scalar.side_effect = [None, mock_user]
_persist_end_user(
sqlite_plugin_engine,
user_id="persisted-user-id",
session_id="daemon-session-uuid",
)
# Act
with app.app_context():
result = get_user("tenant123", "daemon-session-uuid")
# Assert
assert result == mock_user
assert mock_session.scalar.call_count == 2
mock_session.add.assert_not_called()
assert result.id == "persisted-user-id"
with Session(sqlite_plugin_engine) as session:
users = session.scalars(select(EndUser)).all()
assert [user.id for user in users] == ["persisted-user-id"]
@patch("controllers.inner_api.plugin.wraps.select")
@patch("controllers.inner_api.plugin.wraps.EndUser")
@patch("controllers.inner_api.plugin.wraps.sessionmaker")
@patch("controllers.inner_api.plugin.wraps.db")
def test_should_return_existing_anonymous_user_by_session_id(
self, mock_db, mock_sessionmaker, mock_enduser_class, mock_select, app: Flask
self,
sqlite_plugin_engine: Engine,
app: Flask,
):
"""Test returning existing anonymous user by session_id"""
# Arrange
mock_user = MagicMock()
mock_user.session_id = "anonymous_session"
mock_session = MagicMock()
mock_sessionmaker.return_value.begin.return_value.__enter__.return_value = mock_session
mock_session.scalar.return_value = mock_user
mock_query = MagicMock()
mock_select.return_value.where.return_value.limit.return_value = mock_query
_persist_end_user(
sqlite_plugin_engine,
user_id="anonymous-user-id",
session_id="anonymous_session",
is_anonymous=True,
)
# Act
with app.app_context():
result = get_user("tenant123", "anonymous_session")
# Assert
assert result == mock_user
assert result.id == "anonymous-user-id"
@patch("controllers.inner_api.plugin.wraps.select")
@patch("controllers.inner_api.plugin.wraps.EndUser")
@patch("controllers.inner_api.plugin.wraps.sessionmaker")
@patch("controllers.inner_api.plugin.wraps.db")
def test_should_create_new_user_when_not_found(
self, mock_db, mock_sessionmaker, mock_enduser_class, mock_select, app: Flask
self,
sqlite_plugin_engine: Engine,
app: Flask,
):
"""Test creating new user when not found in database"""
# Arrange
mock_session = MagicMock()
mock_sessionmaker.return_value.begin.return_value.__enter__.return_value = mock_session
mock_session.scalar.return_value = None
mock_new_user = MagicMock()
mock_enduser_class.return_value = mock_new_user
mock_query = MagicMock()
mock_select.return_value.where.return_value.limit.return_value = mock_query
# Act
with app.app_context():
result = get_user("tenant123", "user123")
# Assert
assert result == mock_new_user
mock_session.add.assert_called_once()
mock_session.refresh.assert_called_once()
assert result.tenant_id == "tenant123"
assert result.session_id == "user123"
with Session(sqlite_plugin_engine) as session:
persisted_user = session.get(EndUser, result.id)
assert persisted_user is not None
assert persisted_user.session_id == "user123"
@patch("controllers.inner_api.plugin.wraps.select")
@patch("controllers.inner_api.plugin.wraps.EndUser")
@patch("controllers.inner_api.plugin.wraps.sessionmaker")
@patch("controllers.inner_api.plugin.wraps.db")
def test_should_use_default_session_id_when_user_id_none(
self, mock_db, mock_sessionmaker, mock_enduser_class, mock_select, app: Flask
self,
sqlite_plugin_engine: Engine,
app: Flask,
):
"""Test using default session ID when user_id is None"""
# Arrange
mock_user = MagicMock()
mock_session = MagicMock()
mock_sessionmaker.return_value.begin.return_value.__enter__.return_value = mock_session
# When user_id is None, is_anonymous=True, so session.scalar() is used
mock_session.scalar.return_value = mock_user
_persist_end_user(
sqlite_plugin_engine,
user_id="default-user-id",
session_id=DefaultEndUserSessionID.DEFAULT_SESSION_ID,
is_anonymous=True,
)
# Act
with app.app_context():
result = get_user("tenant123", None)
# Assert
assert result == mock_user
assert result.id == "default-user-id"
assert result.session_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID
@patch("controllers.inner_api.plugin.wraps.EndUser")
@patch("controllers.inner_api.plugin.wraps.sessionmaker")
@patch("controllers.inner_api.plugin.wraps.db")
def test_should_raise_error_on_database_exception(self, mock_db, mock_sessionmaker, mock_enduser_class, app: Flask):
def test_should_raise_error_on_database_exception(self, sqlite_plugin_engine: Engine, app: Flask):
"""Test raising ValueError when database operation fails"""
# Arrange
mock_session = MagicMock()
mock_sessionmaker.return_value.begin.return_value.__enter__.return_value = mock_session
mock_session.scalar.side_effect = Exception("Database error")
# Act & Assert
with app.app_context():
with pytest.raises(ValueError, match="user not found"):
def _raise_database_error(*_args, **_kwargs):
raise RuntimeError("Database error")
event.listen(sqlite_plugin_engine, "before_cursor_execute", _raise_database_error)
try:
with app.app_context(), pytest.raises(ValueError, match="user not found"):
get_user("tenant123", "user123")
finally:
event.remove(sqlite_plugin_engine, "before_cursor_execute", _raise_database_error)
class TestGetUserTenant:
"""Test get_user_tenant decorator"""
@patch("controllers.inner_api.plugin.wraps.Tenant")
def test_should_inject_tenant_and_user_models(self, mock_tenant_class, app: Flask, monkeypatch: pytest.MonkeyPatch):
def test_should_inject_tenant_and_user_models(
self,
sqlite_plugin_engine: Engine,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
):
"""Test that decorator injects tenant_model and user_model into kwargs"""
# Arrange
@ -227,24 +241,20 @@ class TestGetUserTenant:
def protected_view(tenant_model, user_model, **kwargs):
return {"tenant": tenant_model, "user": user_model}
mock_tenant = MagicMock()
mock_tenant.id = "tenant123"
mock_user = MagicMock()
mock_user.id = "user456"
_persist_tenant(sqlite_plugin_engine)
_persist_end_user(
sqlite_plugin_engine,
user_id="user456",
session_id="user-session",
)
# Act
with app.test_request_context(json={"tenant_id": "tenant123", "user_id": "user456"}):
monkeypatch.setattr(app, "login_manager", MagicMock(), raising=False)
with patch("controllers.inner_api.plugin.wraps.db.session.get") as mock_get:
with patch("controllers.inner_api.plugin.wraps.get_user") as mock_get_user:
with patch("controllers.inner_api.plugin.wraps.user_logged_in"):
mock_get.return_value = mock_tenant
mock_get_user.return_value = mock_user
result = protected_view()
with patch("controllers.inner_api.plugin.wraps.user_logged_in"):
result = protected_view()
# Assert
assert result["tenant"] == mock_tenant
assert result["user"] == mock_user
assert result["tenant"].id == "tenant123"
assert result["user"].id == "user456"
def test_should_raise_error_when_tenant_id_missing(self, app: Flask):
"""Test that Pydantic ValidationError is raised when tenant_id is missing from payload"""
@ -259,7 +269,7 @@ class TestGetUserTenant:
with pytest.raises(ValidationError):
protected_view()
def test_should_raise_error_when_tenant_not_found(self, app: Flask):
def test_should_raise_error_when_tenant_not_found(self, sqlite_plugin_engine: Engine, app: Flask):
"""Test that ValueError is raised when tenant is not found"""
# Arrange
@ -267,16 +277,15 @@ class TestGetUserTenant:
def protected_view(tenant_model, user_model, **kwargs):
return "success"
# Act & Assert
with app.test_request_context(json={"tenant_id": "nonexistent", "user_id": "user456"}):
with patch("controllers.inner_api.plugin.wraps.db.session.get") as mock_get:
mock_get.return_value = None
with pytest.raises(ValueError, match="tenant not found"):
protected_view()
with pytest.raises(ValueError, match="tenant not found"):
protected_view()
@patch("controllers.inner_api.plugin.wraps.Tenant")
def test_should_use_default_session_id_when_user_id_empty(
self, mock_tenant_class, app: Flask, monkeypatch: pytest.MonkeyPatch
self,
sqlite_plugin_engine: Engine,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
):
"""Test that default session ID is used when user_id is empty string"""
@ -285,26 +294,22 @@ class TestGetUserTenant:
def protected_view(tenant_model, user_model, **kwargs):
return {"tenant": tenant_model, "user": user_model}
mock_tenant = MagicMock()
mock_tenant.id = "tenant123"
mock_user = MagicMock()
_persist_tenant(sqlite_plugin_engine)
_persist_end_user(
sqlite_plugin_engine,
user_id="default-user-id",
session_id=DefaultEndUserSessionID.DEFAULT_SESSION_ID,
is_anonymous=True,
)
# Act - use empty string for user_id to trigger default logic
with app.test_request_context(json={"tenant_id": "tenant123", "user_id": ""}):
monkeypatch.setattr(app, "login_manager", MagicMock(), raising=False)
with patch("controllers.inner_api.plugin.wraps.db.session.get") as mock_get:
with patch("controllers.inner_api.plugin.wraps.get_user") as mock_get_user:
with patch("controllers.inner_api.plugin.wraps.user_logged_in"):
mock_get.return_value = mock_tenant
mock_get_user.return_value = mock_user
result = protected_view()
with patch("controllers.inner_api.plugin.wraps.user_logged_in"):
result = protected_view()
# Assert
assert result["tenant"] == mock_tenant
assert result["user"] == mock_user
from models.model import DefaultEndUserSessionID
mock_get_user.assert_called_once_with("tenant123", DefaultEndUserSessionID.DEFAULT_SESSION_ID)
assert result["tenant"].id == "tenant123"
assert result["user"].id == "default-user-id"
assert result["user"].session_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID
class PluginTestPayload:

View File

@ -2,10 +2,12 @@
Unit tests for inner_api auth decorators
"""
from unittest.mock import MagicMock, patch
from unittest.mock import patch
from uuid import NAMESPACE_URL, uuid5
import pytest
from flask import Flask
from sqlalchemy.orm import Session, sessionmaker
from werkzeug.exceptions import HTTPException
from configs import dify_config
@ -16,9 +18,14 @@ from controllers.inner_api.wraps import (
inner_api_only,
plugin_inner_api_only,
)
from models.enums import EndUserType
from models.model import EndUser
def _stable_uuid(value: str) -> str:
return str(uuid5(NAMESPACE_URL, value))
class TestBillingInnerApiOnly:
"""Test billing_inner_api_only decorator"""
@ -258,7 +265,7 @@ class TestEnterpriseInnerApiUserAuth:
assert result == "no_user"
def test_should_pass_through_when_hmac_signature_invalid(self, app: Flask):
"""Test that request passes through when HMAC signature is invalid"""
"""Invalid HMAC auth passes through without opening a database session."""
# Arrange
@enterprise_inner_api_user_auth
@ -277,7 +284,8 @@ class TestEnterpriseInnerApiUserAuth:
assert result == "no_user"
mock_create_session.assert_not_called()
def test_should_inject_user_when_hmac_signature_valid(self, app: Flask):
@pytest.mark.parametrize("sqlite_session", [(EndUser,)], indirect=True)
def test_should_inject_user_when_hmac_signature_valid(self, app: Flask, sqlite_session: Session):
"""Test that user is injected when HMAC signature is valid"""
# Arrange
from base64 import b64encode
@ -289,19 +297,25 @@ class TestEnterpriseInnerApiUserAuth:
return kwargs.get("user")
# Calculate valid HMAC signature
user_id = "user123"
user_id = _stable_uuid("end-user:user123")
inner_api_key = "valid_key"
data_to_sign = f"DIFY {user_id}"
signature = hmac_new(inner_api_key.encode("utf-8"), data_to_sign.encode("utf-8"), sha1)
valid_signature = b64encode(signature.digest()).decode("utf-8")
# Create mock user
mock_user = MagicMock()
mock_user.id = user_id
mock_session = MagicMock()
mock_session.get.return_value = mock_user
mock_session_context = MagicMock()
mock_session_context.__enter__.return_value = mock_session
end_user = EndUser(
id=user_id,
tenant_id=_stable_uuid("tenant:inner-api"),
type=EndUserType.BROWSER,
name="Inner API User",
session_id="inner-api-session",
)
sqlite_session.add(end_user)
sqlite_session.commit()
database_session_factory = sessionmaker(
bind=sqlite_session.get_bind(),
expire_on_commit=False,
)
# Act
with app.test_request_context(
@ -310,14 +324,15 @@ class TestEnterpriseInnerApiUserAuth:
with patch.object(dify_config, "INNER_API", True):
with patch(
"controllers.inner_api.wraps.session_factory.create_session",
return_value=mock_session_context,
) as mock_create_session:
database_session_factory,
):
result = protected_view()
# Assert
assert result == mock_user
mock_create_session.assert_called_once_with()
mock_session.get.assert_called_once_with(EndUser, user_id)
assert isinstance(result, EndUser)
assert result.id == end_user.id
assert result.tenant_id == end_user.tenant_id
assert result.session_id == "inner-api-session"
class TestPluginInnerApiOnly:

View File

@ -1,14 +1,20 @@
"""Unit tests for runtime credential inner API."""
import inspect
import json
from unittest.mock import MagicMock, patch
import pytest
from flask import Flask
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session
from controllers.inner_api.runtime_credentials import (
EnterpriseRuntimeCredentialsResolve,
InnerRuntimeCredentialsResolvePayload,
)
from models.provider import ProviderCredential
from models.tools import BuiltinToolProvider
def test_runtime_credentials_payload_accepts_items():
@ -32,14 +38,15 @@ def test_runtime_credentials_payload_accepts_items():
@patch("controllers.inner_api.runtime_credentials.encrypter.decrypt_token")
@patch("controllers.inner_api.runtime_credentials.db")
@patch("controllers.inner_api.runtime_credentials.Session")
@patch("controllers.inner_api.runtime_credentials.create_plugin_provider_manager")
@pytest.mark.parametrize("sqlite_session", [(ProviderCredential,)], indirect=True)
def test_runtime_model_credentials_resolve_returns_decrypted_values(
mock_provider_manager_factory,
mock_session_cls,
mock_db,
mock_decrypt_token,
app: Flask,
sqlite_engine: Engine,
sqlite_session: Session,
):
provider_configuration = MagicMock()
provider_configuration.provider.provider_credential_schema.credential_form_schemas = []
@ -52,14 +59,16 @@ def test_runtime_model_credentials_resolve_returns_decrypted_values(
provider_manager.get_configurations.return_value = provider_configurations
mock_provider_manager_factory.return_value = provider_manager
credential = MagicMock()
credential.encrypted_config = '{"openai_api_key":"encrypted","api_base":"https://api.openai.com/v1"}'
session = MagicMock()
session.__enter__.return_value = session
session.__exit__.return_value = False
session.execute.return_value.scalar_one_or_none.return_value = credential
mock_session_cls.return_value = session
mock_db.engine = MagicMock()
credential = ProviderCredential(
tenant_id="tenant-1",
provider_name="langgenius/openai/openai",
credential_name="OpenAI",
encrypted_config='{"openai_api_key":"encrypted","api_base":"https://api.openai.com/v1"}',
)
credential.id = "credential-1"
sqlite_session.add(credential)
sqlite_session.commit()
mock_db.engine = sqlite_engine
mock_decrypt_token.return_value = "sk-test"
handler = EnterpriseRuntimeCredentialsResolve()
@ -110,28 +119,32 @@ def test_runtime_model_credentials_resolve_rejects_unknown_provider(mock_provide
@patch("controllers.inner_api.runtime_credentials.create_provider_encrypter")
@patch("controllers.inner_api.runtime_credentials.ToolProviderCredentialsCache")
@patch("controllers.inner_api.runtime_credentials.db")
@patch("controllers.inner_api.runtime_credentials.Session")
@patch("controllers.inner_api.runtime_credentials.ToolManager")
@pytest.mark.parametrize("sqlite_session", [(BuiltinToolProvider,)], indirect=True)
def test_runtime_tool_credentials_resolve_returns_decrypted_values(
mock_tool_manager,
mock_session_cls,
mock_db,
mock_cache_cls,
mock_create_encrypter,
app: Flask,
sqlite_engine: Engine,
sqlite_session: Session,
):
provider_controller = MagicMock()
provider_controller.get_credentials_schema_by_type.return_value = []
mock_tool_manager.get_builtin_provider.return_value = provider_controller
builtin_provider = MagicMock()
builtin_provider = BuiltinToolProvider(
tenant_id="tenant-1",
user_id="user-1",
provider="langgenius/tavily/tavily",
name="Tavily",
encrypted_credentials=json.dumps({"tavily_api_key": "encrypted"}),
)
builtin_provider.id = "credential-1"
session = MagicMock()
session.__enter__.return_value = session
session.__exit__.return_value = False
session.execute.return_value.scalar_one_or_none.return_value = builtin_provider
mock_session_cls.return_value = session
mock_db.engine = MagicMock()
sqlite_session.add(builtin_provider)
sqlite_session.commit()
mock_db.engine = sqlite_engine
provider_encrypter = MagicMock()
provider_encrypter.decrypt.return_value = {"tavily_api_key": "tvly-secret"}
@ -157,27 +170,34 @@ def test_runtime_tool_credentials_resolve_returns_decrypted_values(
assert body["credentials"][0]["kind"] == "tool"
assert body["credentials"][0]["provider"] == "langgenius/tavily/tavily"
assert body["credentials"][0]["values"]["tavily_api_key"] == "tvly-secret"
compiled = str(session.execute.call_args.args[0].compile(compile_kwargs={"literal_binds": True}))
assert "tool_builtin_providers.provider = 'langgenius/tavily/tavily'" in compiled
provider_encrypter.decrypt.assert_called_once_with({"tavily_api_key": "encrypted"})
@patch("controllers.inner_api.runtime_credentials.db")
@patch("controllers.inner_api.runtime_credentials.Session")
@patch("controllers.inner_api.runtime_credentials.ToolManager")
@pytest.mark.parametrize("sqlite_session", [(BuiltinToolProvider,)], indirect=True)
def test_runtime_tool_credentials_resolve_rejects_unknown_credential(
mock_tool_manager,
mock_session_cls,
mock_db,
app: Flask,
sqlite_engine: Engine,
sqlite_session: Session,
):
mock_tool_manager.get_builtin_provider.return_value = MagicMock()
session = MagicMock()
session.__enter__.return_value = session
session.__exit__.return_value = False
session.execute.return_value.scalar_one_or_none.return_value = None
mock_session_cls.return_value = session
mock_db.engine = MagicMock()
# The requested id exists for another tenant, proving the resolver does not
# expose a credential across workspace boundaries.
builtin_provider = BuiltinToolProvider(
tenant_id="tenant-2",
user_id="user-2",
provider="langgenius/tavily/tavily",
name="Other workspace Tavily",
encrypted_credentials=json.dumps({"tavily_api_key": "encrypted"}),
)
builtin_provider.id = "missing"
sqlite_session.add(builtin_provider)
sqlite_session.commit()
mock_db.engine = sqlite_engine
handler = EnterpriseRuntimeCredentialsResolve()
unwrapped = inspect.unwrap(handler.post)

View File

@ -1,4 +1,9 @@
"""Tests for openapi workflow events reconnect endpoint."""
"""Tests for the OpenAPI workflow-events reconnect endpoint.
The controller constructs a repository session factory, so every case binds
that real SQLAlchemy factory to an isolated SQLite engine. Repository behavior
remains mocked because these tests focus on authorization and SSE responses.
"""
from __future__ import annotations
@ -9,6 +14,8 @@ from unittest.mock import Mock
import pytest
from flask import Flask
from sqlalchemy.engine import Engine
from sqlalchemy.orm import sessionmaker
from werkzeug.exceptions import NotFound
from controllers.openapi.auth.data import AuthData
@ -47,6 +54,11 @@ def _make_workflow_run(
class TestOpenApiWorkflowEventsApi:
@pytest.fixture(autouse=True)
def _bind_sqlite_engine(self, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None:
module = sys.modules["controllers.openapi.workflow_events"]
monkeypatch.setattr(module, "db", SimpleNamespace(engine=sqlite_engine))
def _get_api(self):
from controllers.openapi.workflow_events import OpenApiWorkflowEventsApi
@ -59,8 +71,6 @@ class TestOpenApiWorkflowEventsApi:
factory_mock = Mock()
factory_mock.create_api_workflow_run_repository.return_value = repo_mock
monkeypatch.setattr(module, "DifyAPIRepositoryFactory", factory_mock)
monkeypatch.setattr(module, "sessionmaker", Mock(return_value=object()))
monkeypatch.setattr(module, "db", SimpleNamespace(engine=object()))
api = self._get_api()
from models.model import AppMode
@ -77,6 +87,10 @@ class TestOpenApiWorkflowEventsApi:
auth_data=_make_auth_data(app_model, caller, "account"),
)
session_maker = factory_mock.create_api_workflow_run_repository.call_args.args[0]
assert isinstance(session_maker, sessionmaker)
assert session_maker.kw["bind"] is module.db.engine
def test_not_found_when_run_belongs_to_different_app(
self, app: Flask, bypass_pipeline, monkeypatch: pytest.MonkeyPatch
):
@ -87,8 +101,6 @@ class TestOpenApiWorkflowEventsApi:
factory_mock = Mock()
factory_mock.create_api_workflow_run_repository.return_value = repo_mock
monkeypatch.setattr(module, "DifyAPIRepositoryFactory", factory_mock)
monkeypatch.setattr(module, "sessionmaker", Mock(return_value=object()))
monkeypatch.setattr(module, "db", SimpleNamespace(engine=object()))
api = self._get_api()
from models.model import AppMode
@ -116,8 +128,6 @@ class TestOpenApiWorkflowEventsApi:
factory_mock = Mock()
factory_mock.create_api_workflow_run_repository.return_value = repo_mock
monkeypatch.setattr(module, "DifyAPIRepositoryFactory", factory_mock)
monkeypatch.setattr(module, "sessionmaker", Mock(return_value=object()))
monkeypatch.setattr(module, "db", SimpleNamespace(engine=object()))
snapshot_builder = Mock(return_value=iter([]))
monkeypatch.setattr(module, "build_workflow_event_stream", snapshot_builder)
@ -156,8 +166,6 @@ class TestOpenApiWorkflowEventsApi:
factory_mock = Mock()
factory_mock.create_api_workflow_run_repository.return_value = repo_mock
monkeypatch.setattr(module, "DifyAPIRepositoryFactory", factory_mock)
monkeypatch.setattr(module, "sessionmaker", Mock(return_value=object()))
monkeypatch.setattr(module, "db", SimpleNamespace(engine=object()))
from models.model import AppMode
@ -185,8 +193,6 @@ class TestOpenApiWorkflowEventsApi:
factory_mock = Mock()
factory_mock.create_api_workflow_run_repository.return_value = repo_mock
monkeypatch.setattr(module, "DifyAPIRepositoryFactory", factory_mock)
monkeypatch.setattr(module, "sessionmaker", Mock(return_value=object()))
monkeypatch.setattr(module, "db", SimpleNamespace(engine=object()))
msg_gen_mock = Mock()
msg_gen_mock.retrieve_events.return_value = iter([])
@ -227,8 +233,6 @@ class TestOpenApiWorkflowEventsApi:
factory_mock = Mock()
factory_mock.create_api_workflow_run_repository.return_value = repo_mock
monkeypatch.setattr(module, "DifyAPIRepositoryFactory", factory_mock)
monkeypatch.setattr(module, "sessionmaker", Mock(return_value=object()))
monkeypatch.setattr(module, "db", SimpleNamespace(engine=object()))
finish_response = SimpleNamespace(
event=SimpleNamespace(value="workflow_finished"),

View File

@ -22,6 +22,8 @@ from unittest.mock import Mock, patch
import pytest
from flask import Flask
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session
from werkzeug.exceptions import BadRequest, NotFound
import services
@ -39,19 +41,45 @@ from controllers.service_api.app.conversation import (
ConversationVariableUpdatePayload,
)
from controllers.service_api.app.error import NotChatAppError
from core.app.entities.app_invoke_entities import InvokeFrom
from fields._value_type_serializer import serialize_value_type
from graphon.variables import StringSegment
from graphon.variables.types import SegmentType
from models.model import App, AppMode, EndUser
from models.enums import ConversationFromSource
from models.model import App, AppMode, Conversation, EndUser
from services.conversation_service import ConversationService
from services.errors.conversation import (
ConversationNotExistsError,
ConversationVariableNotExistsError,
ConversationVariableTypeMismatchError,
LastConversationNotExistsError,
)
def _end_user(user_id: str = "end-user-1") -> EndUser:
end_user = EndUser()
end_user.id = user_id
return end_user
def _conversation(
*,
conversation_id: str,
app_id: str = "app-1",
end_user_id: str = "end-user-1",
) -> Conversation:
conversation = Conversation(
app_id=app_id,
mode=AppMode.CHAT,
name="Original Name",
from_source=ConversationFromSource.API,
from_end_user_id=end_user_id,
invoke_from=InvokeFrom.SERVICE_API,
)
conversation.id = conversation_id
conversation.inputs = {}
return conversation
class TestConversationListQuery:
"""Test suite for ConversationListQuery Pydantic model."""
@ -462,23 +490,30 @@ class TestConversationService:
assert hasattr(result, "limit")
assert hasattr(result, "has_more")
@patch.object(ConversationService, "rename")
def test_rename_returns_conversation(self, mock_rename):
@pytest.mark.parametrize("sqlite_session", [(Conversation,)], indirect=True)
def test_rename_returns_conversation(self, sqlite_session: Session):
"""Test rename returns updated conversation."""
mock_conversation = Mock()
mock_conversation.name = "New Name"
mock_rename.return_value = mock_conversation
conversation_id = "00000000-0000-0000-0000-000000000001"
conversation = _conversation(conversation_id=conversation_id)
sqlite_session.add(conversation)
sqlite_session.commit()
app_model = App()
app_model.id = "app-1"
end_user = _end_user()
result = ConversationService.rename(
app_model=Mock(spec=App),
conversation_id="conv_123",
user=Mock(spec=EndUser),
app_model=app_model,
conversation_id=conversation_id,
user=end_user,
name="New Name",
auto_generate=False,
session=Mock(),
session=sqlite_session,
)
assert result.name == "New Name"
sqlite_session.refresh(conversation)
assert conversation.name == "New Name"
class TestConversationPayloadsController:
@ -502,37 +537,29 @@ class TestConversationApiController:
with pytest.raises(NotChatAppError):
handler(api, app_model=app_model, end_user=end_user)
def test_list_last_not_found(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
class _BeginStub:
def __enter__(self):
return SimpleNamespace()
@pytest.mark.parametrize("sqlite_session", [(Conversation,)], indirect=True)
def test_list_last_not_found(
self,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
) -> None:
last_id = "00000000-0000-0000-0000-000000000001"
# The id exists for a different app, proving pagination cannot cross app boundaries.
sqlite_session.add(_conversation(conversation_id=last_id, app_id="other-app"))
sqlite_session.commit()
def __exit__(self, exc_type, exc, tb):
return False
class _SessionMakerStub:
def __init__(self, *args, **kwargs):
pass
def begin(self):
return _BeginStub()
monkeypatch.setattr(
ConversationService,
"pagination_by_last_id",
lambda *_args, **_kwargs: (_ for _ in ()).throw(LastConversationNotExistsError()),
)
conversation_module = sys.modules["controllers.service_api.app.conversation"]
monkeypatch.setattr(conversation_module, "db", SimpleNamespace(engine=object()))
monkeypatch.setattr(conversation_module, "sessionmaker", _SessionMakerStub)
monkeypatch.setattr(conversation_module, "db", SimpleNamespace(engine=sqlite_engine))
api = ConversationApi()
handler = unwrap(api.get)
app_model = SimpleNamespace(mode=AppMode.CHAT)
end_user = SimpleNamespace()
app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT)
end_user = _end_user()
with app.test_request_context(
"/conversations?last_id=00000000-0000-0000-0000-000000000001&limit=20",
f"/conversations?last_id={last_id}&limit=20",
method="GET",
):
with pytest.raises(NotFound):

View File

@ -14,6 +14,8 @@ from unittest.mock import ANY, MagicMock, Mock
import pytest
from flask import Flask
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, sessionmaker
import services.app_generate_service as ags_module
from controllers.service_api.app.workflow_events import WorkflowEventsApi
@ -31,7 +33,7 @@ from core.app.entities.task_entities import (
from core.app.layers.pause_state_persist_layer import WorkflowResumptionContext, _WorkflowGenerateEntityWrapper
from core.workflow.human_input_policy import FormDisposition, HumanInputSurface
from core.workflow.nodes.human_input.entities import ParagraphInputConfig, UserActionConfig
from core.workflow.nodes.human_input.enums import FormInputType
from core.workflow.nodes.human_input.enums import FormInputType, HumanInputFormKind, HumanInputFormStatus
from core.workflow.nodes.human_input.pause_reason import DifyHITLEventType, HumanInputRequired
from core.workflow.system_variables import build_system_variables
from graphon.entities import WorkflowStartReason
@ -39,6 +41,7 @@ from graphon.enums import WorkflowExecutionStatus, WorkflowNodeExecutionStatus
from graphon.runtime import GraphRuntimeState, VariablePool
from models.account import Account
from models.enums import CreatorUserRole
from models.human_input import HumanInputForm
from models.model import AppMode
from models.workflow import WorkflowRun
from repositories.api_workflow_node_execution_repository import WorkflowNodeExecutionSnapshot
@ -66,7 +69,7 @@ class _DummyRateLimit:
return generator
def _mock_repo_for_run(monkeypatch: pytest.MonkeyPatch, workflow_run):
def _mock_repo_for_run(monkeypatch: pytest.MonkeyPatch, workflow_run, sqlite_engine: Engine):
workflow_events_module = sys.modules["controllers.service_api.app.workflow_events"]
repo = SimpleNamespace(get_workflow_run_by_id_and_tenant_id=lambda **_kwargs: workflow_run)
monkeypatch.setattr(
@ -74,10 +77,33 @@ def _mock_repo_for_run(monkeypatch: pytest.MonkeyPatch, workflow_run):
"create_api_workflow_run_repository",
lambda *_args, **_kwargs: repo,
)
monkeypatch.setattr(workflow_events_module, "db", SimpleNamespace(engine=object()))
monkeypatch.setattr(workflow_events_module, "db", SimpleNamespace(engine=sqlite_engine))
return workflow_events_module
def _persist_human_input_form(
sqlite_session: Session,
*,
expiration_time: datetime,
) -> HumanInputForm:
form = HumanInputForm(
id="form-1",
tenant_id="tenant-1",
app_id="app-1",
workflow_run_id="run-1",
conversation_id=None,
form_kind=HumanInputFormKind.RUNTIME,
node_id="node-1",
form_definition=json.dumps({"display_in_ui": True}),
rendered_content="Rendered",
status=HumanInputFormStatus.WAITING,
expiration_time=expiration_time,
)
sqlite_session.add(form)
sqlite_session.commit()
return form
def _build_service_api_pause_converter() -> WorkflowResponseConverter:
application_generate_entity = SimpleNamespace(
inputs={},
@ -257,7 +283,10 @@ def _build_resumption_context(task_id: str) -> WorkflowResumptionContext:
class TestHitlServiceApi:
# Service API event-stream continuation
def test_workflow_events_continue_on_pause_keeps_stream_open(
self, app: Flask, monkeypatch: pytest.MonkeyPatch
self,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
) -> None:
workflow_run = SimpleNamespace(
id="run-1",
@ -266,7 +295,11 @@ class TestHitlServiceApi:
created_by="end-user-1",
finished_at=None,
)
workflow_events_module = _mock_repo_for_run(monkeypatch, workflow_run=workflow_run)
workflow_events_module = _mock_repo_for_run(
monkeypatch,
workflow_run=workflow_run,
sqlite_engine=sqlite_engine,
)
msg_generator = Mock()
msg_generator.retrieve_events.return_value = ["raw-event"]
workflow_generator = Mock()
@ -291,7 +324,10 @@ class TestHitlServiceApi:
workflow_generator.convert_to_event_stream.assert_called_once_with(["raw-event"])
def test_workflow_events_snapshot_continue_on_pause_keeps_pause_open(
self, app: Flask, monkeypatch: pytest.MonkeyPatch
self,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
) -> None:
workflow_run = SimpleNamespace(
id="run-1",
@ -300,7 +336,11 @@ class TestHitlServiceApi:
created_by="end-user-1",
finished_at=None,
)
workflow_events_module = _mock_repo_for_run(monkeypatch, workflow_run=workflow_run)
workflow_events_module = _mock_repo_for_run(
monkeypatch,
workflow_run=workflow_run,
sqlite_engine=sqlite_engine,
)
msg_generator = Mock()
workflow_generator = Mock()
workflow_generator.convert_to_event_stream.return_value = iter(["data: snapshot\n\n"])
@ -331,16 +371,24 @@ class TestHitlServiceApi:
human_input_surface=HumanInputSurface.SERVICE_API,
close_on_pause=False,
)
snapshot_session_maker = snapshot_builder.call_args.kwargs["session_maker"]
assert isinstance(snapshot_session_maker, sessionmaker)
assert snapshot_session_maker.kw["bind"] is sqlite_engine
workflow_generator.convert_to_event_stream.assert_called_once_with(["snapshot-events"])
def test_advanced_chat_blocking_injects_pause_state_config(self, monkeypatch: pytest.MonkeyPatch) -> None:
def test_advanced_chat_blocking_injects_pause_state_config(
self,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
) -> None:
monkeypatch.setattr(ags_module.dify_config, "BILLING_ENABLED", False)
monkeypatch.setattr(ags_module, "RateLimit", _DummyRateLimit)
workflow = MagicMock()
workflow.created_by = "owner-id"
monkeypatch.setattr(AppGenerateService, "_get_workflow", lambda *args, **kwargs: workflow)
monkeypatch.setattr(ags_module.session_factory, "get_session_maker", lambda: "session-maker")
sqlite_session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
monkeypatch.setattr(ags_module.session_factory, "get_session_maker", lambda: sqlite_session_maker)
generator_instance = MagicMock()
generator_instance.generate.return_value = {"result": "advanced-blocking"}
@ -358,20 +406,21 @@ class TestHitlServiceApi:
user = MagicMock()
user.id = "user-id"
result = AppGenerateService.generate(
session=Mock(),
app_model=app_model,
user=user,
args={"workflow_id": None, "query": "hi", "inputs": {}},
invoke_from=InvokeFrom.SERVICE_API,
streaming=False,
)
with sqlite_session_maker() as session:
result = AppGenerateService.generate(
session=session,
app_model=app_model,
user=user,
args={"workflow_id": None, "query": "hi", "inputs": {}},
invoke_from=InvokeFrom.SERVICE_API,
streaming=False,
)
assert result == {"result": "advanced-blocking"}
call_kwargs = generator_instance.generate.call_args.kwargs
assert call_kwargs["streaming"] is False
assert call_kwargs["pause_state_config"] is not None
assert call_kwargs["pause_state_config"].session_factory == "session-maker"
assert call_kwargs["pause_state_config"].session_factory is sqlite_session_maker
assert call_kwargs["pause_state_config"].state_owner_user_id == "owner-id"
# Blocking payload contract
@ -569,7 +618,13 @@ class TestHitlServiceApi:
assert response.data.paused_nodes == ["node-1"]
assert response.data.reasons == [{"TYPE": "human_input_required", "form_id": "form-1", "expiration_time": 1}]
def test_service_api_pause_event_serializes_hitl_reason(self, monkeypatch: pytest.MonkeyPatch) -> None:
@pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True)
def test_service_api_pause_event_serializes_hitl_reason(
self,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
) -> None:
converter = _build_service_api_pause_converter()
converter.workflow_start_to_stream_response(
task_id="task",
@ -578,20 +633,10 @@ class TestHitlServiceApi:
reason=WorkflowStartReason.INITIAL,
)
expiration_time = datetime(2024, 1, 1, tzinfo=UTC)
expiration_time = datetime(2024, 1, 1)
_persist_human_input_form(sqlite_session, expiration_time=expiration_time)
class _FakeSession:
def execute(self, _stmt):
return [("form-1", expiration_time, '{"display_in_ui": true}')]
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
monkeypatch.setattr(workflow_response_converter, "Session", lambda **_: _FakeSession())
monkeypatch.setattr(workflow_response_converter, "db", SimpleNamespace(engine=object()))
monkeypatch.setattr(workflow_response_converter, "db", SimpleNamespace(engine=sqlite_engine))
monkeypatch.setattr(
workflow_response_converter,
"load_form_dispositions_by_form_id",
@ -651,10 +696,18 @@ class TestHitlServiceApi:
assert hi_resp.data.expiration_time == int(expiration_time.timestamp())
# Snapshot payload contract
def test_snapshot_events_include_pause_payload_contract(self, monkeypatch: pytest.MonkeyPatch) -> None:
@pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True)
def test_snapshot_events_include_pause_payload_contract(
self,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
) -> None:
workflow_run = _build_workflow_run(WorkflowExecutionStatus.PAUSED)
snapshot = _build_snapshot(WorkflowNodeExecutionStatus.PAUSED)
resumption_context = _build_resumption_context("task-ctx")
expiration_time = datetime(2024, 1, 1)
_persist_human_input_form(sqlite_session, expiration_time=expiration_time)
monkeypatch.setattr(
"services.workflow_event_snapshot_service.load_form_dispositions_by_form_id",
lambda form_ids, session=None, surface=None: {
@ -662,22 +715,7 @@ class TestHitlServiceApi:
},
)
class _SessionContext:
def __init__(self, session):
self._session = session
def __enter__(self):
return self._session
def __exit__(self, exc_type, exc, tb):
return False
def session_maker() -> _SessionContext:
return _SessionContext(
SimpleNamespace(
execute=lambda _stmt: [("form-1", datetime(2024, 1, 1, tzinfo=UTC), '{"display_in_ui": true}')],
)
)
sqlite_session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
pause_entity = _FakePauseEntity(
pause_id="pause-1",
@ -701,7 +739,7 @@ class TestHitlServiceApi:
message_context=None,
pause_entity=pause_entity,
resumption_context=resumption_context,
session_maker=session_maker,
session_maker=sqlite_session_maker,
)
assert [event["event"] for event in events] == [
@ -713,13 +751,13 @@ class TestHitlServiceApi:
]
assert events[2]["data"]["status"] == WorkflowNodeExecutionStatus.PAUSED.value
assert events[3]["data"]["form_token"] == "wtok"
assert events[3]["data"]["expiration_time"] == int(datetime(2024, 1, 1, tzinfo=UTC).timestamp())
assert events[3]["data"]["expiration_time"] == int(expiration_time.timestamp())
pause_data = events[-1]["data"]
assert pause_data["paused_nodes"] == ["node-1"]
assert pause_data["outputs"] == {"result": "value"}
assert pause_data["reasons"][0]["TYPE"] == "human_input_required"
assert pause_data["reasons"][0]["form_token"] == "wtok"
assert pause_data["reasons"][0]["expiration_time"] == int(datetime(2024, 1, 1, tzinfo=UTC).timestamp())
assert pause_data["reasons"][0]["expiration_time"] == int(expiration_time.timestamp())
assert pause_data["status"] == WorkflowExecutionStatus.PAUSED.value
assert pause_data["created_at"] == int(workflow_run.created_at.timestamp())
assert pause_data["elapsed_time"] == workflow_run.elapsed_time

View File

@ -16,20 +16,20 @@ Focus on:
import json
import sys
import uuid
from dataclasses import dataclass, field
from datetime import UTC, datetime
from inspect import unwrap
from types import SimpleNamespace
from unittest.mock import MagicMock, Mock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import sessionmaker
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, sessionmaker
from werkzeug.exceptions import BadRequest, NotFound
from controllers.service_api.app.error import NotWorkflowAppError, WorkflowVersionExecutionNotAllowedError
from controllers.service_api.app.workflow import (
AppQueueManager,
DifyAPIRepositoryFactory,
GraphEngineManager,
WorkflowAppLogApi,
WorkflowLogQuery,
@ -44,6 +44,7 @@ from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpErr
from core.app.entities.app_invoke_entities import InvokeFrom
from enums.cloud_plan import CloudPlan
from graphon.enums import WorkflowExecutionStatus
from models import Account
from models.enums import CreatorUserRole, WorkflowRunTriggeredFrom
from models.model import App, AppMode, EndUser
from models.workflow import WorkflowAppLog, WorkflowAppLogCreatedFrom, WorkflowRun, WorkflowType
@ -51,58 +52,18 @@ from services.app_generate_service import AppGenerateService
from services.billing_service import BillingService
from services.errors.app import IsDraftWorkflowError, WorkflowNotFoundError
from services.errors.llm import InvokeRateLimitError
from services.workflow_app_service import LogView, LogViewDetails, WorkflowAppService
from services.workflow_app_service import WorkflowAppService
def _default_workflow_inputs() -> dict[str, object]:
return {"input": "value"}
def _default_log_details() -> LogViewDetails:
return {"trigger_metadata": {"node": "answer", "latency": 1.25}}
class _DbSessionStub:
def get(self, *args: object, **kwargs: object) -> None:
return None
@dataclass
class _DbStub:
engine: object = field(default_factory=object)
session: _DbSessionStub = field(default_factory=_DbSessionStub)
@dataclass
class _WorkflowRunRepositoryStub:
run: WorkflowRun | None
def get_workflow_run_by_id(self, *, tenant_id: str, app_id: str, run_id: str) -> WorkflowRun | None:
return self.run if tenant_id and app_id and run_id else None
def get_workflow_run_by_id_without_tenant(self, *, run_id: str) -> WorkflowRun | None:
return self.run if run_id else None
class _BeginStub:
def __enter__(self) -> object:
return object()
def __exit__(self, exc_type: object, exc: object, tb: object) -> bool:
return False
class _SessionMakerStub:
def __init__(self, *args: object, **kwargs: object) -> None:
pass
def begin(self) -> _BeginStub:
return _BeginStub()
def _make_workflow_run(
run_id: str = "run-1",
*,
tenant_id: str = "tenant-1",
app_id: str = "app-1",
workflow_id: str = "wf-1",
inputs: dict[str, object] | None = None,
outputs: dict[str, object] | None = None,
@ -111,8 +72,8 @@ def _make_workflow_run(
) -> WorkflowRun:
return WorkflowRun(
id=run_id,
tenant_id="tenant-1",
app_id="app-1",
tenant_id=tenant_id,
app_id=app_id,
workflow_id=workflow_id,
type=WorkflowType.WORKFLOW,
triggered_from=WorkflowRunTriggeredFrom.APP_RUN,
@ -133,12 +94,17 @@ def _make_workflow_run(
)
def _make_workflow_app_log() -> WorkflowAppLog:
def _make_workflow_app_log(
*,
tenant_id: str = "tenant-1",
app_id: str = "app-1",
workflow_run_id: str = "log-run-1",
) -> WorkflowAppLog:
log = WorkflowAppLog(
tenant_id="tenant-1",
app_id="app-1",
tenant_id=tenant_id,
app_id=app_id,
workflow_id="wf-1",
workflow_run_id="log-run-1",
workflow_run_id=workflow_run_id,
created_from=WorkflowAppLogCreatedFrom.SERVICE_API,
created_by_role=CreatorUserRole.ACCOUNT,
created_by="account-1",
@ -148,16 +114,6 @@ def _make_workflow_app_log() -> WorkflowAppLog:
return log
def _make_workflow_log_page() -> dict[str, object]:
return {
"page": 1,
"limit": 20,
"total": 1,
"has_more": False,
"data": [LogView(_make_workflow_app_log(), _default_log_details())],
}
def _make_app_model(
*,
app_id: str = "app-1",
@ -177,6 +133,43 @@ def _make_end_user(user_id: str = "end-user-1") -> EndUser:
return end_user
def _bind_sqlite_database(
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
) -> None:
"""Bind controller- and model-owned database access to the test engine."""
database = SimpleNamespace(engine=sqlite_engine, session=sqlite_session)
monkeypatch.setattr(sys.modules["controllers.service_api.app.workflow"], "db", database)
monkeypatch.setattr(sys.modules["models.workflow"], "db", database)
def _persist_workflow_log(
sqlite_session: Session,
*,
tenant_id: str,
app_id: str,
) -> None:
workflow_run_id = "log-run-1"
sqlite_session.add_all(
[
_make_workflow_run(
run_id=workflow_run_id,
tenant_id=tenant_id,
app_id=app_id,
created_at=datetime(2026, 1, 1, 1, tzinfo=UTC),
finished_at=datetime(2026, 1, 1, 1, 0, 2, tzinfo=UTC),
),
_make_workflow_app_log(
tenant_id=tenant_id,
app_id=app_id,
workflow_run_id=workflow_run_id,
),
]
)
sqlite_session.commit()
def _expected_workflow_log_pagination_payload() -> dict[str, object]:
return {
"page": 1,
@ -195,16 +188,16 @@ def _expected_workflow_log_pagination_payload() -> dict[str, object]:
"elapsed_time": 0.1,
"total_tokens": 10,
"total_steps": 1,
"created_at": 1767229200,
"finished_at": 1767229202,
"created_at": int(datetime(2026, 1, 1, 1).timestamp()),
"finished_at": int(datetime(2026, 1, 1, 1, 0, 2).timestamp()),
"exceptions_count": 0,
},
"details": {"trigger_metadata": {"node": "answer", "latency": 1.25}},
"details": None,
"created_from": "service-api",
"created_by_role": "account",
"created_by_account": None,
"created_by_end_user": None,
"created_at": 1767229203,
"created_at": int(datetime(2026, 1, 1, 1, 0, 3).timestamp()),
}
],
}
@ -364,15 +357,15 @@ class TestWorkflowAppService:
assert hasattr(WorkflowAppService, "get_paginate_workflow_app_logs")
assert callable(WorkflowAppService.get_paginate_workflow_app_logs)
@patch.object(WorkflowAppService, "get_paginate_workflow_app_logs")
def test_get_paginate_workflow_app_logs_returns_pagination(self, mock_get_logs):
"""Test get_paginate_workflow_app_logs returns paginated result."""
pagination = _make_workflow_log_page()
mock_get_logs.return_value = pagination
@pytest.mark.parametrize("sqlite_session", [(WorkflowAppLog,)], indirect=True)
def test_get_paginate_workflow_app_logs_returns_pagination(self, sqlite_session: Session):
"""Test pagination returns committed logs scoped to the requested app."""
log = _make_workflow_app_log()
sqlite_session.add(log)
sqlite_session.commit()
service = WorkflowAppService()
result = service.get_paginate_workflow_app_logs(
session=Mock(),
session=sqlite_session,
app_model=_make_app_model(),
keyword=None,
status=None,
@ -384,7 +377,11 @@ class TestWorkflowAppService:
created_by_account=None,
)
assert result == pagination
assert result["page"] == 1
assert result["limit"] == 20
assert result["total"] == 1
assert result["has_more"] is False
assert [item.id for item in result["data"]] == [log.id]
class TestWorkflowExecutionStatus:
@ -409,8 +406,9 @@ class TestWorkflowExecutionStatus:
class TestAppGenerateServiceWorkflow:
"""Test AppGenerateService workflow integration."""
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
@patch.object(AppGenerateService, "generate")
def test_generate_accepts_workflow_args(self, mock_generate: MagicMock):
def test_generate_accepts_workflow_args(self, mock_generate: MagicMock, sqlite_session: Session):
"""Test generate accepts workflow-specific args."""
mock_generate.return_value = {"result": "success"}
@ -419,15 +417,17 @@ class TestAppGenerateServiceWorkflow:
user=_make_end_user(),
args={"inputs": {"key": "value"}, "workflow_id": "workflow_123"},
invoke_from=InvokeFrom.SERVICE_API,
session=MagicMock(),
session=sqlite_session,
streaming=False,
)
assert result == {"result": "success"}
mock_generate.assert_called_once()
assert mock_generate.call_args.kwargs["session"] is sqlite_session
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
@patch.object(AppGenerateService, "generate")
def test_generate_raises_workflow_not_found_error(self, mock_generate: MagicMock):
def test_generate_raises_workflow_not_found_error(self, mock_generate: MagicMock, sqlite_session: Session):
"""Test generate raises WorkflowNotFoundError."""
mock_generate.side_effect = WorkflowNotFoundError("Workflow not found")
@ -437,12 +437,13 @@ class TestAppGenerateServiceWorkflow:
user=_make_end_user(),
args={"workflow_id": "invalid_id"},
invoke_from=InvokeFrom.SERVICE_API,
session=MagicMock(),
session=sqlite_session,
streaming=False,
)
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
@patch.object(AppGenerateService, "generate")
def test_generate_raises_is_draft_workflow_error(self, mock_generate: MagicMock):
def test_generate_raises_is_draft_workflow_error(self, mock_generate: MagicMock, sqlite_session: Session):
"""Test generate raises IsDraftWorkflowError."""
mock_generate.side_effect = IsDraftWorkflowError("Workflow is draft")
@ -452,12 +453,13 @@ class TestAppGenerateServiceWorkflow:
user=_make_end_user(),
args={"workflow_id": "draft_workflow"},
invoke_from=InvokeFrom.SERVICE_API,
session=MagicMock(),
session=sqlite_session,
streaming=False,
)
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
@patch.object(AppGenerateService, "generate")
def test_generate_supports_streaming_mode(self, mock_generate: MagicMock):
def test_generate_supports_streaming_mode(self, mock_generate: MagicMock, sqlite_session: Session):
"""Test generate supports streaming response mode."""
mock_stream = Mock()
mock_generate.return_value = mock_stream
@ -467,7 +469,7 @@ class TestAppGenerateServiceWorkflow:
user=_make_end_user(),
args={"inputs": {}, "response_mode": "streaming"},
invoke_from=InvokeFrom.SERVICE_API,
session=MagicMock(),
session=sqlite_session,
streaming=True,
)
@ -499,19 +501,23 @@ class TestWorkflowRunRepository:
assert hasattr(DifyAPIRepositoryFactory, "create_api_workflow_run_repository")
@patch("repositories.factory.DifyAPIRepositoryFactory.create_api_workflow_run_repository")
def test_workflow_run_repository_get_by_id(self, mock_factory):
"""Test workflow run repository get_workflow_run_by_id method."""
@pytest.mark.parametrize("sqlite_session", [(WorkflowRun,)], indirect=True)
def test_workflow_run_repository_get_by_id(self, sqlite_engine: Engine, sqlite_session: Session):
"""Test repository lookup against committed tenant-scoped state."""
run = _make_workflow_run(run_id=str(uuid.uuid4()))
mock_factory.return_value = _WorkflowRunRepositoryStub(run=run)
sqlite_session.add(run)
sqlite_session.commit()
from repositories.factory import DifyAPIRepositoryFactory
repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(sessionmaker())
repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(
sessionmaker(bind=sqlite_engine, expire_on_commit=False)
)
result = repo.get_workflow_run_by_id(tenant_id="tenant_123", app_id="app_456", run_id="run_789")
result = repo.get_workflow_run_by_id(tenant_id="tenant-1", app_id="app-1", run_id=run.id)
assert result == run
assert result is not None
assert result.id == run.id
assert repo.get_workflow_run_by_id(tenant_id="other-tenant", app_id="app-1", run_id=run.id) is None
class TestWorkflowRunDetailApi:
@ -524,16 +530,17 @@ class TestWorkflowRunDetailApi:
with pytest.raises(NotWorkflowAppError):
handler(api, app_model=app_model, workflow_run_id="run")
def test_success(self, monkeypatch: pytest.MonkeyPatch) -> None:
run = _make_workflow_run(run_id="run")
repo = _WorkflowRunRepositoryStub(run=run)
workflow_module = sys.modules["controllers.service_api.app.workflow"]
monkeypatch.setattr(workflow_module, "db", _DbStub())
monkeypatch.setattr(
DifyAPIRepositoryFactory,
"create_api_workflow_run_repository",
lambda *_args, **_kwargs: repo,
)
@pytest.mark.parametrize("sqlite_session", [(WorkflowRun,)], indirect=True)
def test_success(
self,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
) -> None:
run = _make_workflow_run(run_id="run", tenant_id="t1", app_id="a1")
sqlite_session.add(run)
sqlite_session.commit()
_bind_sqlite_database(monkeypatch, sqlite_engine, sqlite_session)
api = WorkflowRunDetailApi()
handler = unwrap(api.get)
@ -546,7 +553,8 @@ class TestWorkflowRunDetailApi:
class TestWorkflowRunApi:
def test_not_workflow_app(self, app: Flask) -> None:
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_not_workflow_app(self, app: Flask, sqlite_session: Session) -> None:
api = WorkflowRunApi()
handler = unwrap(api.post)
app_model = _make_app_model(mode=AppMode.CHAT)
@ -554,9 +562,10 @@ class TestWorkflowRunApi:
with app.test_request_context("/workflows/run", method="POST", json={"inputs": {}}):
with pytest.raises(NotWorkflowAppError):
handler(api, session=Mock(), app_model=app_model, end_user=end_user)
handler(api, session=sqlite_session, app_model=app_model, end_user=end_user)
def test_rate_limit(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_rate_limit(self, app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
monkeypatch.setattr(
AppGenerateService,
"generate",
@ -570,7 +579,7 @@ class TestWorkflowRunApi:
with app.test_request_context("/workflows/run", method="POST", json={"inputs": {}}):
with pytest.raises(InvokeRateLimitHttpError):
handler(api, session=Mock(), app_model=app_model, end_user=end_user)
handler(api, session=sqlite_session, app_model=app_model, end_user=end_user)
def test_sandbox_billing_does_not_gate_default_workflow_run(
self, app: Flask, monkeypatch: pytest.MonkeyPatch
@ -680,7 +689,8 @@ class TestWorkflowRunByIdApi:
else:
billing_get_info.assert_not_called()
def test_not_found(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_not_found(self, app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
workflow_module = sys.modules["controllers.service_api.app.workflow"]
monkeypatch.setattr(workflow_module.dify_config, "BILLING_ENABLED", False)
monkeypatch.setattr(
@ -696,9 +706,10 @@ class TestWorkflowRunByIdApi:
with app.test_request_context("/workflows/1/run", method="POST", json={"inputs": {}}):
with pytest.raises(NotFound):
handler(api, session=Mock(), app_model=app_model, end_user=end_user, workflow_id="w1")
handler(api, session=sqlite_session, app_model=app_model, end_user=end_user, workflow_id="w1")
def test_draft_workflow(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_draft_workflow(self, app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
workflow_module = sys.modules["controllers.service_api.app.workflow"]
monkeypatch.setattr(workflow_module.dify_config, "BILLING_ENABLED", False)
monkeypatch.setattr(
@ -714,7 +725,7 @@ class TestWorkflowRunByIdApi:
with app.test_request_context("/workflows/1/run", method="POST", json={"inputs": {}}):
with pytest.raises(BadRequest):
handler(api, session=Mock(), app_model=app_model, end_user=end_user, workflow_id="w1")
handler(api, session=sqlite_session, app_model=app_model, end_user=end_user, workflow_id="w1")
class TestWorkflowTaskStopApi:
@ -748,28 +759,16 @@ class TestWorkflowTaskStopApi:
class TestWorkflowAppLogApi:
def test_success(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
workflow_module = sys.modules["controllers.service_api.app.workflow"]
workflow_model_module = sys.modules["models.workflow"]
monkeypatch.setattr(workflow_module, "db", _DbStub())
monkeypatch.setattr(workflow_model_module, "db", _DbStub())
monkeypatch.setattr(workflow_module, "sessionmaker", _SessionMakerStub)
monkeypatch.setattr(
WorkflowAppService,
"get_paginate_workflow_app_logs",
lambda *_args, **_kwargs: _make_workflow_log_page(),
)
monkeypatch.setattr(
DifyAPIRepositoryFactory,
"create_api_workflow_run_repository",
lambda *_args, **_kwargs: _WorkflowRunRepositoryStub(
run=_make_workflow_run(
run_id="log-run-1",
created_at=datetime(2026, 1, 1, 1, tzinfo=UTC),
finished_at=datetime(2026, 1, 1, 1, 0, 2, tzinfo=UTC),
)
),
)
@pytest.mark.parametrize("sqlite_session", [(WorkflowRun, WorkflowAppLog, Account)], indirect=True)
def test_success(
self,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
) -> None:
_persist_workflow_log(sqlite_session, tenant_id="tenant-1", app_id="a1")
_bind_sqlite_database(monkeypatch, sqlite_engine, sqlite_session)
api = WorkflowAppLogApi()
handler = unwrap(api.get)
@ -803,18 +802,24 @@ class TestWorkflowRunDetailApiGet:
and we call the unwrapped method directly in tests.
"""
@patch("controllers.service_api.app.workflow.DifyAPIRepositoryFactory")
@patch("controllers.service_api.app.workflow.db")
@pytest.mark.parametrize("sqlite_session", [(WorkflowRun,)], indirect=True)
def test_get_workflow_run_success(
self,
mock_db,
mock_repo_factory,
app: Flask,
workflow_app: App,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
):
"""Test successful workflow run detail retrieval."""
run = _make_workflow_run(run_id="run-1")
mock_repo_factory.create_api_workflow_run_repository.return_value = _WorkflowRunRepositoryStub(run=run)
run = _make_workflow_run(
run_id="run-1",
tenant_id=workflow_app.tenant_id,
app_id=workflow_app.id,
)
sqlite_session.add(run)
sqlite_session.commit()
_bind_sqlite_database(monkeypatch, sqlite_engine, sqlite_session)
from controllers.service_api.app.workflow import WorkflowRunDetailApi
@ -834,13 +839,12 @@ class TestWorkflowRunDetailApiGet:
"error": None,
"total_steps": 1,
"total_tokens": 10,
"created_at": 1767225600,
"finished_at": 1767225600,
"created_at": int(datetime(2026, 1, 1).timestamp()),
"finished_at": int(datetime(2026, 1, 1).timestamp()),
"elapsed_time": 0.1,
}
@patch("controllers.service_api.app.workflow.db")
def test_get_workflow_run_wrong_app_mode(self, mock_db, app: Flask):
def test_get_workflow_run_wrong_app_mode(self, app: Flask):
"""Test NotWorkflowAppError when app mode is not workflow or advanced_chat."""
from controllers.service_api.app.workflow import WorkflowRunDetailApi
@ -902,46 +906,23 @@ class TestWorkflowAppLogApiGet:
``get`` is wrapped by ``@validate_app_token``.
"""
@patch("controllers.service_api.app.workflow.WorkflowAppService")
@patch("controllers.service_api.app.workflow.db")
@pytest.mark.parametrize("sqlite_session", [(WorkflowRun, WorkflowAppLog, Account)], indirect=True)
def test_get_workflow_logs_success(
self,
mock_db,
mock_wf_svc_cls,
app: Flask,
workflow_app: App,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
):
"""Test successful workflow log retrieval."""
mock_svc_instance = Mock()
mock_svc_instance.get_paginate_workflow_app_logs.return_value = _make_workflow_log_page()
mock_wf_svc_cls.return_value = mock_svc_instance
mock_repo = _WorkflowRunRepositoryStub(
run=_make_workflow_run(
run_id="log-run-1",
created_at=datetime(2026, 1, 1, 1, tzinfo=UTC),
finished_at=datetime(2026, 1, 1, 1, 0, 2, tzinfo=UTC),
)
)
# Mock sessionmaker(...).begin() context manager
mock_db.engine = object()
mock_db.session.get.return_value = None
_persist_workflow_log(sqlite_session, tenant_id=workflow_app.tenant_id, app_id=workflow_app.id)
_bind_sqlite_database(monkeypatch, sqlite_engine, sqlite_session)
from controllers.service_api.app.workflow import WorkflowAppLogApi
with app.test_request_context(
"/workflows/logs?page=1&limit=20",
method="GET",
):
with (
patch("controllers.service_api.app.workflow.sessionmaker", _SessionMakerStub),
patch("models.workflow.db", _DbStub()),
patch(
"repositories.factory.DifyAPIRepositoryFactory.create_api_workflow_run_repository",
return_value=mock_repo,
),
):
api = WorkflowAppLogApi()
result = unwrap(api.get)(api, app_model=workflow_app)
with app.test_request_context("/workflows/logs?page=1&limit=20", method="GET"):
api = WorkflowAppLogApi()
result = unwrap(api.get)(api, app_model=workflow_app)
assert result == _expected_workflow_log_pagination_payload()

View File

@ -7,18 +7,57 @@ Service API controller tests.
"""
import uuid
from collections.abc import Iterator
from dataclasses import dataclass
from unittest.mock import Mock
import pytest
from flask import Flask
from sqlalchemy import Engine
from sqlalchemy.orm import Session
from core.rag.index_processor.constant.index_type import IndexStructureType
from models.account import TenantStatus
from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole, TenantStatus
from models.base import TypeBase
from models.model import App, AppMode, EndUser
from tests.unit_tests.conftest import (
setup_mock_dataset_owner_execute_result,
setup_mock_tenant_owner_execute_result,
)
@dataclass(frozen=True)
class ServiceApiIdentity:
"""Persisted owner identity for service-API authentication tests."""
session: Session
tenant: Tenant
account: Account
membership: TenantAccountJoin
@pytest.fixture
def service_api_identity(sqlite_engine: Engine) -> Iterator[ServiceApiIdentity]:
"""Yield an isolated SQLite session with a real active tenant owner."""
TypeBase.metadata.create_all(
sqlite_engine,
tables=[Account.__table__, Tenant.__table__, TenantAccountJoin.__table__],
)
with Session(sqlite_engine, expire_on_commit=False) as session:
tenant = Tenant(name="Service API Workspace")
tenant.id = str(uuid.uuid4())
account = Account(name="Service API Owner", email=f"owner-{tenant.id}@example.com")
account.id = str(uuid.uuid4())
membership = TenantAccountJoin(
tenant_id=tenant.id,
account_id=account.id,
role=TenantAccountRole.OWNER,
)
account._current_tenant = tenant
session.add_all([tenant, account, membership])
session.commit()
yield ServiceApiIdentity(
session=session,
tenant=tenant,
account=account,
membership=membership,
)
@pytest.fixture
@ -110,40 +149,6 @@ def mock_dataset_api_token(mock_tenant_id):
return token
class AuthenticationMocker:
"""
Helper class to set up common authentication mocking patterns.
Usage:
auth_mocker = AuthenticationMocker()
with auth_mocker.mock_app_auth(mock_api_token, mock_app_model, mock_tenant):
# Test code here
"""
@staticmethod
def setup_db_queries(mock_db, mock_app, mock_tenant, mock_account=None):
"""Configure mock_db to return app and tenant via session.get()."""
mock_db.session.get.side_effect = [mock_app, mock_tenant]
if mock_account:
setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account)
@staticmethod
def setup_dataset_auth(mock_db, mock_tenant, mock_account):
"""Configure mock_db for dataset token authentication."""
mock_ta = Mock()
mock_ta.account_id = mock_account.id
setup_mock_dataset_owner_execute_result(mock_db, mock_tenant, mock_ta)
mock_db.session.get.return_value = mock_account
@pytest.fixture
def auth_mocker():
"""Provide an AuthenticationMocker instance."""
return AuthenticationMocker()
@pytest.fixture
def mock_dataset():
"""Create a mock Dataset model."""

View File

@ -0,0 +1,55 @@
"""State-based checks for shared service-API authentication fixtures."""
from uuid import uuid4
from sqlalchemy import select
from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole
from tests.unit_tests.conftest import (
persist_service_api_dataset_owner,
persist_service_api_tenant_owner,
)
from tests.unit_tests.controllers.service_api.conftest import ServiceApiIdentity
def test_service_api_identity_persists_tenant_scoped_owner(service_api_identity: ServiceApiIdentity) -> None:
identity = service_api_identity
owner_row = identity.session.execute(
select(Tenant, Account)
.join(TenantAccountJoin, Tenant.id == TenantAccountJoin.tenant_id)
.join(Account, TenantAccountJoin.account_id == Account.id)
.where(
Tenant.id == identity.tenant.id,
TenantAccountJoin.role == TenantAccountRole.OWNER,
)
).one()
assert owner_row == (identity.tenant, identity.account)
assert identity.account.current_tenant is identity.tenant
def test_shared_helpers_persist_real_app_and_dataset_owner_rows(service_api_identity: ServiceApiIdentity) -> None:
session = service_api_identity.session
app_tenant = Tenant(name="App Workspace")
app_tenant.id = str(uuid4())
app_owner = Account(name="App Owner", email=f"app-owner-{app_tenant.id}@example.com")
app_owner.id = str(uuid4())
app_membership = persist_service_api_tenant_owner(session, app_tenant, app_owner)
dataset_tenant = Tenant(name="Dataset Workspace")
dataset_tenant.id = str(uuid4())
dataset_membership = TenantAccountJoin(
tenant_id=dataset_tenant.id,
account_id=service_api_identity.account.id,
role=TenantAccountRole.OWNER,
)
persist_service_api_dataset_owner(session, dataset_tenant, dataset_membership)
assert session.get(TenantAccountJoin, app_membership.id) is app_membership
assert session.execute(
select(Tenant, TenantAccountJoin)
.join(TenantAccountJoin, Tenant.id == TenantAccountJoin.tenant_id)
.where(Tenant.id == dataset_tenant.id)
).one() == (dataset_tenant, dataset_membership)

View File

@ -574,6 +574,27 @@ def test_console_account_avatar_query_param_renders_as_query(monkeypatch: pytest
assert params["avatar"]["required"] is True
def test_console_agent_debug_conversation_refresh_body_is_optional(monkeypatch: pytest.MonkeyPatch):
from configs import dify_config
from controllers.console import bp as console_bp
monkeypatch.setattr(dify_config, "SWAGGER_UI_ENABLED", True)
app = Flask(__name__)
app.config["TESTING"] = True
app.register_blueprint(console_bp)
payload = app.test_client().get("/console/api/openapi.json").get_json()
operation = payload["paths"]["/agent/{agent_id}/debug-conversation/refresh"]["post"]
request_body = operation["requestBody"]
assert request_body["required"] is False
assert request_body["content"]["application/json"]["schema"] == {
"$ref": "#/components/schemas/AgentDebugConversationRefreshPayload"
}
assert "AgentDebugConversationRefreshPayload" in payload["components"]["schemas"]
def test_console_member_invite_documents_bad_request_response(monkeypatch: pytest.MonkeyPatch):
from configs import dify_config
from controllers.console import bp as console_bp

View File

@ -0,0 +1,67 @@
from unittest.mock import MagicMock, patch
from configs import dify_config
from controllers.web import site as site_module
from extensions.storage.storage_type import StorageType
from models.model import IconType, Site
def test_build_site_icon_url_uses_s3_presigned_url() -> None:
site = MagicMock(spec=Site)
site.icon_type = IconType.IMAGE
site.icon = "11111111-1111-4111-8111-111111111111"
with (
patch.object(dify_config, "EDITION", "CLOUD"),
patch.object(dify_config, "STORAGE_TYPE", StorageType.S3),
patch.object(site_module, "db") as mock_db,
patch.object(site_module, "FileService") as mock_file_service,
patch.object(site_module, "build_icon_url") as mock_build_icon_url,
):
mock_file_service.return_value.get_file_presigned_url.return_value = (
"https://s3.example.com/icon.png?signature=test"
)
result = site_module._build_site_icon_url(site=site, tenant_id="tenant-id")
assert result == "https://s3.example.com/icon.png?signature=test"
mock_file_service.assert_called_once_with(mock_db.engine)
mock_file_service.return_value.get_file_presigned_url.assert_called_once_with(
file_id="11111111-1111-4111-8111-111111111111",
tenant_id="tenant-id",
)
mock_build_icon_url.assert_not_called()
def test_build_site_icon_url_keeps_preview_url_for_self_hosted_s3() -> None:
site = MagicMock(spec=Site)
site.icon_type = IconType.IMAGE
site.icon = "11111111-1111-4111-8111-111111111111"
with (
patch.object(dify_config, "EDITION", "SELF_HOSTED"),
patch.object(dify_config, "STORAGE_TYPE", StorageType.S3),
patch.object(site_module, "FileService") as mock_file_service,
patch.object(site_module, "build_icon_url", return_value="https://api.example.com/files/icon/file-preview"),
):
result = site_module._build_site_icon_url(site=site, tenant_id="tenant-id")
assert result == "https://api.example.com/files/icon/file-preview"
mock_file_service.assert_not_called()
def test_build_site_icon_url_keeps_preview_url_for_non_s3_storage() -> None:
site = MagicMock(spec=Site)
site.icon_type = IconType.IMAGE
site.icon = "11111111-1111-4111-8111-111111111111"
with (
patch.object(dify_config, "EDITION", "CLOUD"),
patch.object(dify_config, "STORAGE_TYPE", StorageType.LOCAL),
patch.object(site_module, "FileService") as mock_file_service,
patch.object(site_module, "build_icon_url", return_value="https://api.example.com/files/icon/file-preview"),
):
result = site_module._build_site_icon_url(site=site, tenant_id="tenant-id")
assert result == "https://api.example.com/files/icon/file-preview"
mock_file_service.assert_not_called()

View File

@ -1,459 +1,120 @@
"""Test conversation variable handling in AdvancedChatAppRunner."""
"""SQLite-backed conversation-variable synchronization tests for AdvancedChatAppRunner."""
from unittest.mock import MagicMock, patch
from uuid import uuid4
from unittest.mock import MagicMock
import pytest
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.app.apps.advanced_chat import app_runner as app_runner_module
from core.app.apps.advanced_chat.app_runner import AdvancedChatAppRunner
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom
from factories import variable_factory
from graphon.variables import SegmentType
from models import ConversationVariable, Workflow
from models import ConversationVariable
MINIMAL_GRAPH = {
"nodes": [
APP_ID = "11111111-1111-1111-1111-111111111111"
CONVERSATION_ID = "22222222-2222-2222-2222-222222222222"
OTHER_CONVERSATION_ID = "22222222-2222-2222-2222-222222222223"
VAR_1_ID = "33333333-3333-3333-3333-333333333333"
VAR_2_ID = "33333333-3333-3333-3333-333333333334"
def _variable(variable_id: str, name: str, value: str):
return variable_factory.build_conversation_variable_from_mapping(
{
"id": "start",
"data": {
"type": "start",
"title": "Start",
},
"id": variable_id,
"name": name,
"value_type": SegmentType.STRING,
"value": value,
}
],
"edges": [],
}
)
def _patch_create_session(mock_session: MagicMock):
session_context = MagicMock()
session_context.__enter__.return_value = mock_session
session_context.__exit__.return_value = False
mock_session.begin.return_value.__enter__.return_value = mock_session
mock_session.begin.return_value.__exit__.return_value = False
return patch("core.app.apps.advanced_chat.app_runner.create_session", return_value=session_context)
def _runner(workflow_variables: list[object]) -> AdvancedChatAppRunner:
workflow = MagicMock()
workflow.conversation_variables = workflow_variables
conversation = MagicMock(app_id=APP_ID, id=CONVERSATION_ID)
return AdvancedChatAppRunner(
application_generate_entity=MagicMock(),
queue_manager=MagicMock(),
conversation=conversation,
message=MagicMock(),
dialogue_count=1,
variable_loader=MagicMock(),
workflow=workflow,
system_user_id="44444444-4444-4444-4444-444444444444",
app=MagicMock(),
workflow_execution_repository=MagicMock(),
workflow_node_execution_repository=MagicMock(),
)
class TestAdvancedChatAppRunnerConversationVariables:
"""Test that AdvancedChatAppRunner correctly handles conversation variables."""
def _bind_runner_sessions(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
engine = sqlite_session.get_bind()
monkeypatch.setattr(
app_runner_module,
"create_session",
lambda: Session(engine, expire_on_commit=False),
)
def test_missing_conversation_variables_are_added(self):
"""Test that new conversation variables added to workflow are created for existing conversations."""
# Setup
app_id = str(uuid4())
conversation_id = str(uuid4())
workflow_id = str(uuid4())
# Create workflow with two conversation variables
workflow_vars = [
variable_factory.build_conversation_variable_from_mapping(
{
"id": "var1",
"name": "existing_var",
"value_type": SegmentType.STRING,
"value": "default1",
}
),
variable_factory.build_conversation_variable_from_mapping(
{
"id": "var2",
"name": "new_var",
"value_type": SegmentType.STRING,
"value": "default2",
}
),
]
# Mock workflow with conversation variables
mock_workflow = MagicMock(spec=Workflow)
mock_workflow.conversation_variables = workflow_vars
mock_workflow.tenant_id = str(uuid4())
mock_workflow.app_id = app_id
mock_workflow.id = workflow_id
mock_workflow.type = "chat"
mock_workflow.graph_dict = MINIMAL_GRAPH
mock_workflow.environment_variables = []
# Create existing conversation variable (only var1 exists in DB)
existing_db_var = MagicMock(spec=ConversationVariable)
existing_db_var.id = "var1"
existing_db_var.app_id = app_id
existing_db_var.conversation_id = conversation_id
existing_db_var.to_variable = MagicMock(return_value=workflow_vars[0])
# Mock conversation and message
mock_conversation = MagicMock()
mock_conversation.app_id = app_id
mock_conversation.id = conversation_id
mock_message = MagicMock()
mock_message.id = str(uuid4())
# Mock app config
mock_app_config = MagicMock()
mock_app_config.app_id = app_id
mock_app_config.workflow_id = workflow_id
mock_app_config.tenant_id = str(uuid4())
# Mock app generate entity
mock_app_generate_entity = MagicMock(spec=AdvancedChatAppGenerateEntity)
mock_app_generate_entity.app_config = mock_app_config
mock_app_generate_entity.inputs = {}
mock_app_generate_entity.query = "test query"
mock_app_generate_entity.files = []
mock_app_generate_entity.user_id = str(uuid4())
mock_app_generate_entity.invoke_from = InvokeFrom.SERVICE_API
mock_app_generate_entity.workflow_run_id = str(uuid4())
mock_app_generate_entity.task_id = str(uuid4())
mock_app_generate_entity.call_depth = 0
mock_app_generate_entity.single_iteration_run = None
mock_app_generate_entity.single_loop_run = None
mock_app_generate_entity.extras = {}
mock_app_generate_entity.trace_manager = None
# Create runner
runner = AdvancedChatAppRunner(
application_generate_entity=mock_app_generate_entity,
queue_manager=MagicMock(),
conversation=mock_conversation,
message=mock_message,
dialogue_count=1,
variable_loader=MagicMock(),
workflow=mock_workflow,
system_user_id=str(uuid4()),
app=MagicMock(),
workflow_execution_repository=MagicMock(),
workflow_node_execution_repository=MagicMock(),
def _persist_variable(session: Session, *, variable: object, conversation_id: str = CONVERSATION_ID) -> None:
session.add(
ConversationVariable.from_variable(
app_id=APP_ID,
conversation_id=conversation_id,
variable=variable,
)
)
session.commit()
# Mock database session
mock_session = MagicMock(spec=Session)
# First query returns only existing variable
mock_scalars_result = MagicMock()
mock_scalars_result.all.return_value = [existing_db_var]
mock_session.scalars.return_value = mock_scalars_result
@pytest.mark.parametrize("sqlite_session", [(ConversationVariable,)], indirect=True)
def test_missing_conversation_variables_are_added(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
existing_variable = _variable(VAR_1_ID, "existing_var", "default1")
new_variable = _variable(VAR_2_ID, "new_var", "default2")
_persist_variable(sqlite_session, variable=existing_variable)
_persist_variable(sqlite_session, variable=new_variable, conversation_id=OTHER_CONVERSATION_ID)
_bind_runner_sessions(monkeypatch, sqlite_session)
# Track what gets added to session
added_items = []
variables = _runner([existing_variable, new_variable])._initialize_conversation_variables()
def track_add_all(items):
added_items.extend(items)
assert [variable.id for variable in variables] == [VAR_1_ID, VAR_2_ID]
persisted = sqlite_session.scalars(
select(ConversationVariable)
.where(ConversationVariable.conversation_id == CONVERSATION_ID)
.order_by(ConversationVariable.id)
).all()
assert [variable.id for variable in persisted] == [VAR_1_ID, VAR_2_ID]
mock_session.add_all.side_effect = track_add_all
# Patch the necessary components
with (
_patch_create_session(mock_session),
patch("core.app.apps.advanced_chat.app_runner.select") as mock_select,
patch.object(runner, "_init_graph") as mock_init_graph,
patch.object(
runner,
"handle_input_moderation",
return_value=(False, mock_app_generate_entity.inputs, mock_app_generate_entity.query),
),
patch.object(runner, "handle_annotation_reply", return_value=False),
patch("core.app.apps.advanced_chat.app_runner.WorkflowEntry") as mock_workflow_entry_class,
patch("core.app.apps.advanced_chat.app_runner.GraphRuntimeState") as mock_graph_runtime_state_class,
patch("core.app.apps.advanced_chat.app_runner.redis_client") as mock_redis_client,
patch("core.app.apps.advanced_chat.app_runner.RedisChannel") as mock_redis_channel_class,
):
# Mock GraphRuntimeState to accept the variable pool
mock_graph_runtime_state_class.return_value = MagicMock()
@pytest.mark.parametrize("sqlite_session", [(ConversationVariable,)], indirect=True)
def test_no_variables_creates_all(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
workflow_variables = [
_variable(VAR_1_ID, "var1", "default1"),
_variable(VAR_2_ID, "var2", "default2"),
]
_bind_runner_sessions(monkeypatch, sqlite_session)
# Mock graph initialization
mock_init_graph.return_value = MagicMock()
variables = _runner(workflow_variables)._initialize_conversation_variables()
# Mock workflow entry
mock_workflow_entry = MagicMock()
mock_workflow_entry.run.return_value = iter([]) # Empty generator
mock_workflow_entry_class.return_value = mock_workflow_entry
assert [variable.id for variable in variables] == [VAR_1_ID, VAR_2_ID]
persisted = sqlite_session.scalars(select(ConversationVariable).order_by(ConversationVariable.id)).all()
assert [variable.id for variable in persisted] == [VAR_1_ID, VAR_2_ID]
# Run the method
runner.run()
# Verify that the missing variable was added
assert len(added_items) == 1, "Should have added exactly one missing variable"
@pytest.mark.parametrize("sqlite_session", [(ConversationVariable,)], indirect=True)
def test_all_variables_exist_no_changes(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
workflow_variables = [
_variable(VAR_1_ID, "var1", "default1"),
_variable(VAR_2_ID, "var2", "default2"),
]
for variable in workflow_variables:
_persist_variable(sqlite_session, variable=variable)
_bind_runner_sessions(monkeypatch, sqlite_session)
# Check that the added item is the missing variable (var2)
added_var = added_items[0]
assert hasattr(added_var, "id"), "Added item should be a ConversationVariable"
# Note: Since we're mocking ConversationVariable.from_variable,
# we can't directly check the id, but we can verify add_all was called
assert mock_session.add_all.called, "Session add_all should have been called"
variables = _runner(workflow_variables)._initialize_conversation_variables()
def test_no_variables_creates_all(self):
"""Test that all conversation variables are created when none exist in DB."""
# Setup
app_id = str(uuid4())
conversation_id = str(uuid4())
workflow_id = str(uuid4())
# Create workflow with conversation variables
workflow_vars = [
variable_factory.build_conversation_variable_from_mapping(
{
"id": "var1",
"name": "var1",
"value_type": SegmentType.STRING,
"value": "default1",
}
),
variable_factory.build_conversation_variable_from_mapping(
{
"id": "var2",
"name": "var2",
"value_type": SegmentType.STRING,
"value": "default2",
}
),
]
# Mock workflow
mock_workflow = MagicMock(spec=Workflow)
mock_workflow.conversation_variables = workflow_vars
mock_workflow.tenant_id = str(uuid4())
mock_workflow.app_id = app_id
mock_workflow.id = workflow_id
mock_workflow.type = "chat"
mock_workflow.graph_dict = MINIMAL_GRAPH
mock_workflow.environment_variables = []
# Mock conversation and message
mock_conversation = MagicMock()
mock_conversation.app_id = app_id
mock_conversation.id = conversation_id
mock_message = MagicMock()
mock_message.id = str(uuid4())
# Mock app config
mock_app_config = MagicMock()
mock_app_config.app_id = app_id
mock_app_config.workflow_id = workflow_id
mock_app_config.tenant_id = str(uuid4())
# Mock app generate entity
mock_app_generate_entity = MagicMock(spec=AdvancedChatAppGenerateEntity)
mock_app_generate_entity.app_config = mock_app_config
mock_app_generate_entity.inputs = {}
mock_app_generate_entity.query = "test query"
mock_app_generate_entity.files = []
mock_app_generate_entity.user_id = str(uuid4())
mock_app_generate_entity.invoke_from = InvokeFrom.SERVICE_API
mock_app_generate_entity.workflow_run_id = str(uuid4())
mock_app_generate_entity.task_id = str(uuid4())
mock_app_generate_entity.call_depth = 0
mock_app_generate_entity.single_iteration_run = None
mock_app_generate_entity.single_loop_run = None
mock_app_generate_entity.extras = {}
mock_app_generate_entity.trace_manager = None
# Create runner
runner = AdvancedChatAppRunner(
application_generate_entity=mock_app_generate_entity,
queue_manager=MagicMock(),
conversation=mock_conversation,
message=mock_message,
dialogue_count=1,
variable_loader=MagicMock(),
workflow=mock_workflow,
system_user_id=str(uuid4()),
app=MagicMock(),
workflow_execution_repository=MagicMock(),
workflow_node_execution_repository=MagicMock(),
)
# Mock database session
mock_session = MagicMock(spec=Session)
# Query returns empty list (no existing variables)
mock_scalars_result = MagicMock()
mock_scalars_result.all.return_value = []
mock_session.scalars.return_value = mock_scalars_result
# Track what gets added to session
added_items = []
def track_add_all(items):
added_items.extend(items)
mock_session.add_all.side_effect = track_add_all
# Patch the necessary components
with (
_patch_create_session(mock_session),
patch("core.app.apps.advanced_chat.app_runner.select") as mock_select,
patch.object(runner, "_init_graph") as mock_init_graph,
patch.object(
runner,
"handle_input_moderation",
return_value=(False, mock_app_generate_entity.inputs, mock_app_generate_entity.query),
),
patch.object(runner, "handle_annotation_reply", return_value=False),
patch("core.app.apps.advanced_chat.app_runner.WorkflowEntry") as mock_workflow_entry_class,
patch("core.app.apps.advanced_chat.app_runner.GraphRuntimeState") as mock_graph_runtime_state_class,
patch("core.app.apps.advanced_chat.app_runner.ConversationVariable") as mock_conv_var_class,
patch("core.app.apps.advanced_chat.app_runner.redis_client") as mock_redis_client,
patch("core.app.apps.advanced_chat.app_runner.RedisChannel") as mock_redis_channel_class,
):
# Mock ConversationVariable.from_variable to return mock objects
mock_conv_vars = []
for var in workflow_vars:
mock_cv = MagicMock()
mock_cv.id = var.id
mock_cv.to_variable.return_value = var
mock_conv_vars.append(mock_cv)
mock_conv_var_class.from_variable.side_effect = mock_conv_vars
# Mock GraphRuntimeState to accept the variable pool
mock_graph_runtime_state_class.return_value = MagicMock()
# Mock graph initialization
mock_init_graph.return_value = MagicMock()
# Mock workflow entry
mock_workflow_entry = MagicMock()
mock_workflow_entry.run.return_value = iter([]) # Empty generator
mock_workflow_entry_class.return_value = mock_workflow_entry
# Run the method
runner.run()
# Verify that all variables were created
assert len(added_items) == 2, "Should have added both variables"
assert mock_session.add_all.called, "Session add_all should have been called"
def test_all_variables_exist_no_changes(self):
"""Test that no changes are made when all variables already exist in DB."""
# Setup
app_id = str(uuid4())
conversation_id = str(uuid4())
workflow_id = str(uuid4())
# Create workflow with conversation variables
workflow_vars = [
variable_factory.build_conversation_variable_from_mapping(
{
"id": "var1",
"name": "var1",
"value_type": SegmentType.STRING,
"value": "default1",
}
),
variable_factory.build_conversation_variable_from_mapping(
{
"id": "var2",
"name": "var2",
"value_type": SegmentType.STRING,
"value": "default2",
}
),
]
# Mock workflow
mock_workflow = MagicMock(spec=Workflow)
mock_workflow.conversation_variables = workflow_vars
mock_workflow.tenant_id = str(uuid4())
mock_workflow.app_id = app_id
mock_workflow.id = workflow_id
mock_workflow.type = "chat"
mock_workflow.graph_dict = MINIMAL_GRAPH
mock_workflow.environment_variables = []
# Create existing conversation variables (both exist in DB)
existing_db_vars = []
for var in workflow_vars:
db_var = MagicMock(spec=ConversationVariable)
db_var.id = var.id
db_var.app_id = app_id
db_var.conversation_id = conversation_id
db_var.to_variable = MagicMock(return_value=var)
existing_db_vars.append(db_var)
# Mock conversation and message
mock_conversation = MagicMock()
mock_conversation.app_id = app_id
mock_conversation.id = conversation_id
mock_message = MagicMock()
mock_message.id = str(uuid4())
# Mock app config
mock_app_config = MagicMock()
mock_app_config.app_id = app_id
mock_app_config.workflow_id = workflow_id
mock_app_config.tenant_id = str(uuid4())
# Mock app generate entity
mock_app_generate_entity = MagicMock(spec=AdvancedChatAppGenerateEntity)
mock_app_generate_entity.app_config = mock_app_config
mock_app_generate_entity.inputs = {}
mock_app_generate_entity.query = "test query"
mock_app_generate_entity.files = []
mock_app_generate_entity.user_id = str(uuid4())
mock_app_generate_entity.invoke_from = InvokeFrom.SERVICE_API
mock_app_generate_entity.workflow_run_id = str(uuid4())
mock_app_generate_entity.task_id = str(uuid4())
mock_app_generate_entity.call_depth = 0
mock_app_generate_entity.single_iteration_run = None
mock_app_generate_entity.single_loop_run = None
mock_app_generate_entity.extras = {}
mock_app_generate_entity.trace_manager = None
# Create runner
runner = AdvancedChatAppRunner(
application_generate_entity=mock_app_generate_entity,
queue_manager=MagicMock(),
conversation=mock_conversation,
message=mock_message,
dialogue_count=1,
variable_loader=MagicMock(),
workflow=mock_workflow,
system_user_id=str(uuid4()),
app=MagicMock(),
workflow_execution_repository=MagicMock(),
workflow_node_execution_repository=MagicMock(),
)
# Mock database session
mock_session = MagicMock(spec=Session)
# Query returns all existing variables
mock_scalars_result = MagicMock()
mock_scalars_result.all.return_value = existing_db_vars
mock_session.scalars.return_value = mock_scalars_result
# Patch the necessary components
with (
_patch_create_session(mock_session),
patch("core.app.apps.advanced_chat.app_runner.select") as mock_select,
patch.object(runner, "_init_graph") as mock_init_graph,
patch.object(
runner,
"handle_input_moderation",
return_value=(False, mock_app_generate_entity.inputs, mock_app_generate_entity.query),
),
patch.object(runner, "handle_annotation_reply", return_value=False),
patch("core.app.apps.advanced_chat.app_runner.WorkflowEntry") as mock_workflow_entry_class,
patch("core.app.apps.advanced_chat.app_runner.GraphRuntimeState") as mock_graph_runtime_state_class,
patch("core.app.apps.advanced_chat.app_runner.redis_client") as mock_redis_client,
patch("core.app.apps.advanced_chat.app_runner.RedisChannel") as mock_redis_channel_class,
):
# Mock GraphRuntimeState to accept the variable pool
mock_graph_runtime_state_class.return_value = MagicMock()
# Mock graph initialization
mock_init_graph.return_value = MagicMock()
# Mock workflow entry
mock_workflow_entry = MagicMock()
mock_workflow_entry.run.return_value = iter([]) # Empty generator
mock_workflow_entry_class.return_value = mock_workflow_entry
# Run the method
runner.run()
# Verify that no variables were added
assert not mock_session.add_all.called, "Session add_all should not have been called"
assert [variable.id for variable in variables] == [VAR_1_ID, VAR_2_ID]
persisted = sqlite_session.scalars(select(ConversationVariable)).all()
assert len(persisted) == 2

View File

@ -14,7 +14,8 @@ import pytest
from core.app.apps.agent_app.app_generator import AgentAppGenerator, AgentAppGeneratorError, AgentAppNotPublishedError
from core.app.entities.app_invoke_entities import InvokeFrom
from models.agent import AgentConfigDraftType, AgentSource
from models.agent import AgentConfigDraft, AgentConfigDraftType, AgentScope, AgentSource
from models.agent_config_entities import AgentSoulConfig
_SOUL_DICT = {
"model": {
@ -95,7 +96,7 @@ class TestResolveDebugDraft:
created_by="creator-1",
updated_by="updater-1",
)
session = _FakeScalarSession([None, SimpleNamespace(id="agent-1"), _snapshot()])
session = _FakeScalarSession([None, _snapshot()])
draft = AgentAppGenerator._resolve_debug_draft(
tenant_id="t1",
@ -110,6 +111,77 @@ class TestResolveDebugDraft:
assert session.added == [draft]
assert session.flush_count == 1
def test_stale_workflow_only_shared_draft_is_rebased_to_active_snapshot(self):
agent = SimpleNamespace(
id="agent-1",
scope=AgentScope.WORKFLOW_ONLY,
active_config_snapshot_id="snap-2",
created_by="creator-1",
updated_by="updater-1",
)
draft = AgentConfigDraft(
id="draft-1",
tenant_id="t1",
agent_id="agent-1",
draft_type=AgentConfigDraftType.DRAFT,
account_id=None,
draft_owner_key="",
base_snapshot_id="snap-1",
config_snapshot=AgentSoulConfig.model_validate({"prompt": {"system_prompt": "old"}}),
)
active_snapshot = SimpleNamespace(
id="snap-2",
config_snapshot_dict={"prompt": {"system_prompt": "new"}},
)
session = _FakeScalarSession([draft, active_snapshot])
resolved = AgentAppGenerator._resolve_debug_draft(
tenant_id="t1",
agent=agent,
draft_type=None,
account_id=None,
session=session,
)
assert resolved is draft
assert resolved.id == "draft-1"
assert resolved.base_snapshot_id == "snap-2"
assert resolved.config_snapshot_dict["prompt"]["system_prompt"] == "new"
assert session.flush_count == 1
def test_build_draft_is_not_rebased_to_active_snapshot(self):
agent = SimpleNamespace(
id="agent-1",
scope=AgentScope.WORKFLOW_ONLY,
active_config_snapshot_id="snap-2",
created_by="creator-1",
updated_by="updater-1",
)
draft = AgentConfigDraft(
id="build-draft-1",
tenant_id="t1",
agent_id="agent-1",
draft_type=AgentConfigDraftType.DEBUG_BUILD,
account_id="account-1",
draft_owner_key="account-1",
base_snapshot_id="snap-1",
config_snapshot=AgentSoulConfig.model_validate({"prompt": {"system_prompt": "build edit"}}),
)
session = _FakeScalarSession([draft])
resolved = AgentAppGenerator._resolve_debug_draft(
tenant_id="t1",
agent=agent,
draft_type=AgentConfigDraftType.DEBUG_BUILD.value,
account_id="account-1",
session=session,
)
assert resolved is draft
assert resolved.base_snapshot_id == "snap-1"
assert resolved.config_snapshot_dict["prompt"]["system_prompt"] == "build edit"
assert session.flush_count == 0
class TestResolveAgent:
def test_success_chains_to_resolve_by_id(self):
@ -185,9 +257,12 @@ class TestResolveAgent:
def test_unpublished_imported_agent_remains_available_to_debugger(self):
bound_agent = SimpleNamespace(
id="agent-1",
scope=AgentScope.ROSTER,
source=AgentSource.IMPORTED,
active_config_snapshot_id="snap-1",
active_config_is_published=False,
created_by="creator-1",
updated_by="updater-1",
)
draft = SimpleNamespace(id="draft-1", draft_type="draft", config_snapshot_dict=_SOUL_DICT)
session = _FakeScalarSession([bound_agent, draft])

View File

@ -2,10 +2,10 @@ from __future__ import annotations
import logging
from contextlib import nullcontext
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from sqlalchemy import func, select
from sqlalchemy.orm import Session
from core.app.app_config.entities import (
AdvancedChatMessageEntity,
@ -15,7 +15,13 @@ from core.app.app_config.entities import (
)
from core.app.apps.base_app_runner import AppRunner
from core.app.apps.exc import GenerateTaskStoppedError
from core.app.entities.app_invoke_entities import InvokeFrom
from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager
from core.app.entities.app_invoke_entities import (
AppGenerateEntity,
EasyUIBasedAppGenerateEntity,
InvokeFrom,
ModelConfigWithCredentialsEntity,
)
from core.app.entities.queue_entities import (
QueueAgentMessageEvent,
QueueLLMChunkEvent,
@ -29,9 +35,9 @@ from graphon.model_runtime.entities.message_entities import (
PromptMessageRole,
TextPromptMessageContent,
)
from graphon.model_runtime.entities.model_entities import ModelPropertyKey
from graphon.model_runtime.entities.model_entities import AIModelEntity, ModelPropertyKey
from graphon.model_runtime.errors.invoke import InvokeBadRequestError
from models.model import AppMode
from models.model import App, AppMode, Message, MessageFile
class _DummyParameterRule:
@ -40,13 +46,29 @@ class _DummyParameterRule:
self.use_template = use_template
class _QueueRecorder:
def __init__(self) -> None:
self.events: list[object] = []
class _TokenCountingModel:
token_count: int
def publish(self, event, pub_from):
_ = pub_from
self.events.append(event)
def __init__(self, token_count: int) -> None:
self.token_count = token_count
def get_llm_num_tokens(self, messages: list[AssistantPromptMessage]) -> int:
return self.token_count
def _queue_manager() -> MessageBasedAppQueueManager:
return MessageBasedAppQueueManager(
task_id="task-id",
user_id="user-id",
invoke_from=InvokeFrom.SERVICE_API,
conversation_id="conversation-id",
app_mode=AppMode.CHAT.value,
message_id="message-id",
)
def _published_events(queue_manager: MessageBasedAppQueueManager) -> list[object]:
return [message.event for message in queue_manager.listen()]
class _ClosableStream:
@ -70,11 +92,11 @@ class TestAppRunner:
def test_recalc_llm_max_tokens_updates_parameters(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
model_schema = SimpleNamespace(
model_schema = AIModelEntity.model_construct(
model_properties={ModelPropertyKey.CONTEXT_SIZE: 100},
parameter_rules=[_DummyParameterRule("max_tokens")],
)
model_config = SimpleNamespace(
model_config = ModelConfigWithCredentialsEntity.model_construct(
provider_model_bundle=object(),
model="mock",
model_schema=model_schema,
@ -83,7 +105,7 @@ class TestAppRunner:
monkeypatch.setattr(
"core.app.apps.base_app_runner.ModelInstance",
lambda provider_model_bundle, model: SimpleNamespace(get_llm_num_tokens=lambda messages: 80),
lambda provider_model_bundle, model: _TokenCountingModel(80),
)
runner.recalc_llm_max_tokens(model_config, prompt_messages=[AssistantPromptMessage(content="hi")])
@ -93,11 +115,11 @@ class TestAppRunner:
def test_recalc_llm_max_tokens_returns_minus_one_when_no_context(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
model_schema = SimpleNamespace(
model_schema = AIModelEntity.model_construct(
model_properties={},
parameter_rules=[_DummyParameterRule("max_tokens")],
)
model_config = SimpleNamespace(
model_config = ModelConfigWithCredentialsEntity.model_construct(
provider_model_bundle=object(),
model="mock",
model_schema=model_schema,
@ -106,17 +128,16 @@ class TestAppRunner:
monkeypatch.setattr(
"core.app.apps.base_app_runner.ModelInstance",
lambda provider_model_bundle, model: SimpleNamespace(get_llm_num_tokens=lambda messages: 10),
lambda provider_model_bundle, model: _TokenCountingModel(10),
)
assert runner.recalc_llm_max_tokens(model_config, prompt_messages=[]) == -1
def test_direct_output_streaming_publishes_chunks_and_end(self, monkeypatch: pytest.MonkeyPatch):
def test_direct_output_streaming_publishes_chunks_and_end(self):
runner = AppRunner()
queue = _QueueRecorder()
app_generate_entity = SimpleNamespace(model_conf=SimpleNamespace(model="mock"), stream=True)
monkeypatch.setattr("core.app.apps.base_app_runner.time.sleep", lambda _: None)
queue = _queue_manager()
model_config = ModelConfigWithCredentialsEntity.model_construct(model="mock")
app_generate_entity = EasyUIBasedAppGenerateEntity.model_construct(model_conf=model_config, stream=True)
runner.direct_output(
queue_manager=queue,
@ -126,12 +147,13 @@ class TestAppRunner:
stream=True,
)
assert any(isinstance(event, QueueLLMChunkEvent) for event in queue.events)
assert isinstance(queue.events[-1], QueueMessageEndEvent)
events = _published_events(queue)
assert any(isinstance(event, QueueLLMChunkEvent) for event in events)
assert isinstance(events[-1], QueueMessageEndEvent)
def test_handle_invoke_result_direct_publishes_end_event(self):
runner = AppRunner()
queue = _QueueRecorder()
queue = _queue_manager()
llm_result = LLMResult(
model="mock",
prompt_messages=[],
@ -145,11 +167,11 @@ class TestAppRunner:
stream=False,
)
assert isinstance(queue.events[-1], QueueMessageEndEvent)
assert isinstance(_published_events(queue)[-1], QueueMessageEndEvent)
def test_handle_invoke_result_invalid_type_raises(self):
runner = AppRunner()
queue = _QueueRecorder()
queue = _queue_manager()
with pytest.raises(NotImplementedError):
runner._handle_invoke_result(
@ -160,7 +182,7 @@ class TestAppRunner:
def test_organize_prompt_messages_simple_template(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
model_config = SimpleNamespace(mode="chat", stop=["STOP"])
model_config = ModelConfigWithCredentialsEntity.model_construct(mode="chat", stop=["STOP"])
prompt_template_entity = PromptTemplateEntity(
prompt_type=PromptTemplateEntity.PromptType.SIMPLE,
simple_prompt_template="hello",
@ -172,7 +194,7 @@ class TestAppRunner:
)
prompt_messages, stop = runner.organize_prompt_messages(
app_record=SimpleNamespace(mode=AppMode.CHAT.value),
app_record=App(mode=AppMode.CHAT.value),
model_config=model_config,
prompt_template_entity=prompt_template_entity,
inputs={},
@ -185,7 +207,7 @@ class TestAppRunner:
def test_organize_prompt_messages_advanced_completion_template(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
model_config = SimpleNamespace(mode="completion", stop=["<END>"])
model_config = ModelConfigWithCredentialsEntity.model_construct(mode="completion", stop=["<END>"])
captured: dict[str, object] = {}
prompt_template_entity = PromptTemplateEntity(
prompt_type=PromptTemplateEntity.PromptType.ADVANCED,
@ -202,7 +224,7 @@ class TestAppRunner:
monkeypatch.setattr("core.app.apps.base_app_runner.AdvancedPromptTransform.get_prompt", _fake_advanced_prompt)
prompt_messages, stop = runner.organize_prompt_messages(
app_record=SimpleNamespace(mode=AppMode.CHAT.value),
app_record=App(mode=AppMode.CHAT.value),
model_config=model_config,
prompt_template_entity=prompt_template_entity,
inputs={},
@ -218,7 +240,7 @@ class TestAppRunner:
def test_organize_prompt_messages_advanced_chat_template(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
model_config = SimpleNamespace(mode="chat", stop=["<END>"])
model_config = ModelConfigWithCredentialsEntity.model_construct(mode="chat", stop=["<END>"])
captured: dict[str, object] = {}
prompt_template_entity = PromptTemplateEntity(
prompt_type=PromptTemplateEntity.PromptType.ADVANCED,
@ -237,7 +259,7 @@ class TestAppRunner:
monkeypatch.setattr("core.app.apps.base_app_runner.AdvancedPromptTransform.get_prompt", _fake_advanced_prompt)
prompt_messages, stop = runner.organize_prompt_messages(
app_record=SimpleNamespace(mode=AppMode.CHAT.value),
app_record=App(mode=AppMode.CHAT.value),
model_config=model_config,
prompt_template_entity=prompt_template_entity,
inputs={},
@ -254,8 +276,8 @@ class TestAppRunner:
with pytest.raises(InvokeBadRequestError, match="Advanced completion prompt template is required"):
runner.organize_prompt_messages(
app_record=SimpleNamespace(mode=AppMode.CHAT.value),
model_config=SimpleNamespace(mode="completion", stop=[]),
app_record=App(mode=AppMode.CHAT.value),
model_config=ModelConfigWithCredentialsEntity.model_construct(mode="completion", stop=[]),
prompt_template_entity=PromptTemplateEntity(prompt_type=PromptTemplateEntity.PromptType.ADVANCED),
inputs={},
files=[],
@ -263,18 +285,16 @@ class TestAppRunner:
with pytest.raises(InvokeBadRequestError, match="Advanced chat prompt template is required"):
runner.organize_prompt_messages(
app_record=SimpleNamespace(mode=AppMode.CHAT.value),
model_config=SimpleNamespace(mode="chat", stop=[]),
app_record=App(mode=AppMode.CHAT.value),
model_config=ModelConfigWithCredentialsEntity.model_construct(mode="chat", stop=[]),
prompt_template_entity=PromptTemplateEntity(prompt_type=PromptTemplateEntity.PromptType.ADVANCED),
inputs={},
files=[],
)
def test_handle_invoke_result_stream_routes_chunks_and_builds_message(
self, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
):
def test_handle_invoke_result_stream_routes_chunks_and_builds_message(self, caplog: pytest.LogCaptureFixture):
runner = AppRunner()
queue = _QueueRecorder()
queue = _queue_manager()
image_content = ImagePromptMessageContent(
url="https://example.com/image.png", format="png", mime_type="image/png"
@ -286,11 +306,9 @@ class TestAppRunner:
prompt_messages=[AssistantPromptMessage(content="prompt")],
delta=LLMResultChunkDelta(
index=0,
message=AssistantPromptMessage.model_construct(
message=AssistantPromptMessage(
content=[
"a",
TextPromptMessageContent(data="b"),
SimpleNamespace(data="c"),
TextPromptMessageContent(data="abc"),
image_content,
]
),
@ -305,21 +323,25 @@ class TestAppRunner:
agent=False,
)
assert isinstance(queue.events[0], QueueLLMChunkEvent)
assert isinstance(queue.events[-1], QueueMessageEndEvent)
assert queue.events[-1].llm_result.message.content == "abc"
events = _published_events(queue)
assert isinstance(events[0], QueueLLMChunkEvent)
assert isinstance(events[-1], QueueMessageEndEvent)
assert events[-1].llm_result.message.content == "abc"
assert "Received multimodal output but missing required parameters" in caplog.messages
def test_handle_invoke_result_stream_agent_mode_handles_multimodal_errors(
self, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
):
runner = AppRunner()
queue = _QueueRecorder()
queue = _queue_manager()
def raise_multimodal_error(**kwargs):
raise RuntimeError("failed to save image")
monkeypatch.setattr(
runner,
"_handle_multimodal_image_content",
MagicMock(side_effect=RuntimeError("failed to save image")),
raise_multimodal_error,
)
usage = LLMUsage.empty_usage()
@ -353,22 +375,37 @@ class TestAppRunner:
tenant_id="tenant-id",
)
assert isinstance(queue.events[0], QueueAgentMessageEvent)
assert isinstance(queue.events[-1], QueueMessageEndEvent)
assert queue.events[-1].llm_result.usage == usage
events = _published_events(queue)
assert isinstance(events[0], QueueAgentMessageEvent)
assert isinstance(events[-1], QueueMessageEndEvent)
assert events[-1].llm_result.usage == usage
assert "Failed to handle multimodal image output" in caplog.messages
def test_handle_invoke_result_stream_commits_message_file_before_publish(self, monkeypatch: pytest.MonkeyPatch):
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_handle_invoke_result_stream_commits_message_file_before_publish(
self,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
):
runner = AppRunner()
runner._handle_multimodal_image_content = MagicMock(return_value="message-file-1")
session = MagicMock()
monkeypatch.setattr(
runner,
"_handle_multimodal_image_content",
lambda **kwargs: "message-file-1",
)
events: list[str] = []
session.commit.side_effect = lambda: events.append("commit")
original_commit = sqlite_session.commit
def commit():
events.append("commit")
original_commit()
monkeypatch.setattr(sqlite_session, "commit", commit)
monkeypatch.setattr(
"core.app.apps.base_app_runner.session_factory.create_session",
lambda: nullcontext(session),
lambda: nullcontext(sqlite_session),
)
queue = _QueueRecorder()
queue = _queue_manager()
original_publish = queue.publish
def publish(event, pub_from):
@ -407,7 +444,7 @@ class TestAppRunner:
assert events == ["commit", "publish"]
def test_handle_invoke_result_stream_closes_generator_when_stopped(self):
def test_handle_invoke_result_stream_closes_generator_when_stopped(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
chunk = LLMResultChunk(
model="stream-model",
@ -416,9 +453,8 @@ class TestAppRunner:
)
stream = _ClosableStream([chunk])
queue_manager = SimpleNamespace(
publish=MagicMock(side_effect=GenerateTaskStoppedError("stopped")),
)
queue_manager = _queue_manager()
monkeypatch.setattr(queue_manager, "_is_stopped", lambda: True)
with pytest.raises(GenerateTaskStoppedError):
runner._handle_invoke_result_stream(
@ -429,7 +465,11 @@ class TestAppRunner:
assert stream.closed is True
def test_handle_multimodal_image_content_fallback_return_branch(self, monkeypatch: pytest.MonkeyPatch):
@pytest.mark.parametrize("sqlite_session", [(MessageFile,)], indirect=True)
def test_handle_multimodal_image_content_fallback_return_branch(
self,
sqlite_session: Session,
):
runner = AppRunner()
class _ToggleBool:
@ -442,19 +482,17 @@ class TestAppRunner:
self._index += 1
return value
content = SimpleNamespace(
# The fallback is reachable only when the fields change truthiness between the guard and branch checks.
content = ImagePromptMessageContent.model_construct(
url=_ToggleBool([False, False]),
base64_data=_ToggleBool([True, False]),
mime_type="image/png",
)
db_session = SimpleNamespace(add=MagicMock(), flush=MagicMock(), refresh=MagicMock())
monkeypatch.setattr("core.app.apps.base_app_runner.ToolFileManager", lambda: MagicMock())
queue_manager = SimpleNamespace(invoke_from=InvokeFrom.SERVICE_API, publish=MagicMock())
queue_manager = _queue_manager()
runner._handle_multimodal_image_content(
session=db_session,
session=sqlite_session,
content=content,
message_id="message-id",
user_id="user-id",
@ -462,20 +500,20 @@ class TestAppRunner:
queue_manager=queue_manager,
)
db_session.add.assert_not_called()
queue_manager.publish.assert_not_called()
message_file_count = sqlite_session.scalar(select(func.count()).select_from(MessageFile))
assert message_file_count == 0
def test_check_hosting_moderation_direct_output_called(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
queue = _QueueRecorder()
app_generate_entity = SimpleNamespace(stream=False)
queue = _queue_manager()
app_generate_entity = EasyUIBasedAppGenerateEntity.model_construct(stream=False)
direct_output_calls: list[dict[str, object]] = []
monkeypatch.setattr(
"core.app.apps.base_app_runner.HostingModerationFeature.check",
lambda self, application_generate_entity, prompt_messages: True,
)
direct_output = MagicMock()
monkeypatch.setattr(runner, "direct_output", direct_output)
monkeypatch.setattr(runner, "direct_output", lambda **kwargs: direct_output_calls.append(kwargs))
result = runner.check_hosting_moderation(
application_generate_entity=app_generate_entity,
@ -484,7 +522,7 @@ class TestAppRunner:
)
assert result is True
assert direct_output.called
assert len(direct_output_calls) == 1
def test_fill_in_inputs_from_external_data_tools(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
@ -509,7 +547,7 @@ class TestAppRunner:
"core.app.apps.base_app_runner.InputModeration.check",
lambda self, app_id, tenant_id, app_config, inputs, query, message_id, trace_manager: (True, {}, ""),
)
app_generate_entity = SimpleNamespace(app_config=SimpleNamespace(), trace_manager=None)
app_generate_entity = AppGenerateEntity.model_construct(app_config=None, trace_manager=None)
result = runner.moderation_for_inputs(
app_id="app",
@ -522,7 +560,12 @@ class TestAppRunner:
assert result == (True, {}, "")
def test_query_app_annotations_to_reply(self, monkeypatch: pytest.MonkeyPatch):
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_query_app_annotations_to_reply(
self,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
):
runner = AppRunner()
monkeypatch.setattr(
"core.app.apps.base_app_runner.AnnotationReplyFeature.query",
@ -530,12 +573,12 @@ class TestAppRunner:
)
response = runner.query_app_annotations_to_reply(
app_record=SimpleNamespace(),
message=SimpleNamespace(),
app_record=App(),
message=Message(),
query="hello",
user_id="user",
invoke_from=InvokeFrom.WEB_APP,
session=MagicMock(),
session=sqlite_session,
)
assert response == "reply"

View File

@ -7,7 +7,7 @@ from pytest_mock import MockerFixture
from core.plugin.endpoint.exc import EndpointSetupFailedError
from core.plugin.entities.plugin_daemon import PluginDaemonInnerError
from core.plugin.impl.base import PLUGIN_DAEMON_MAX_PATH_LENGTH, BasePluginClient
from core.plugin.impl.exc import PluginLLMPollingUnsupportedError
from core.plugin.impl.exc import PluginLLMPollingUnsupportedError, PluginRuntimeError
from core.trigger.errors import (
EventIgnoreError,
TriggerInvokeError,
@ -175,3 +175,25 @@ class TestBasePluginClientImpl:
with pytest.raises(PluginLLMPollingUnsupportedError):
client._handle_plugin_daemon_error("PluginInvokeError", message)
def test_handle_plugin_daemon_error_maps_runtime_error_to_typed_exception(self):
client = BasePluginClient()
lambda_request_id = "45664803-3d3c-4d4f-93fe-e3b19e43092b"
message = json.dumps(
{
"error_type": PluginRuntimeError.__name__,
"message": (
"Plugin runtime request failed: Runtime.ExitError: "
f"RequestId: {lambda_request_id} Error: Runtime exited with error: exit status 1"
),
"args": {"request_id": lambda_request_id, "status_code": 200},
}
)
with pytest.raises(PluginRuntimeError) as exc_info:
client._handle_plugin_daemon_error("PluginInvokeError", message)
assert exc_info.value.description == (
"Plugin runtime request failed: Runtime.ExitError: Runtime exited with error: exit status 1"
)
assert exc_info.value.lambda_request_id == lambda_request_id

View File

@ -125,6 +125,9 @@ class TestPluginParameterEntities:
parameter = PluginParameter(name="p", label=self._label(), options="invalid") # type: ignore[arg-type]
assert parameter.options == []
def test_plugin_parameter_excludes_tool_specific_multiple_declaration(self):
assert "multiple" not in PluginParameter.model_fields
@pytest.mark.parametrize(
("parameter_type", "expected"),
[

View File

@ -5,6 +5,8 @@ from dataclasses import dataclass
from typing import Any, cast
from unittest.mock import MagicMock
import pytest
from core.app.entities.app_invoke_entities import InvokeFrom
from core.tools.__base.tool import Tool
from core.tools.__base.tool_runtime import ToolRuntime
@ -33,6 +35,7 @@ class DummyParameter:
options: list[Any] | None = None
llm_description: str | None = None
input_schema: dict[str, Any] | None = None
multiple: bool = False
class DummyTool(Tool):
@ -129,6 +132,26 @@ def test_invoke_supports_single_message_and_parameter_casting():
}
def test_invoke_preserves_multiple_select_values():
tool = _build_tool()
parameter = ToolParameter.get_simple_instance(
name="choice",
llm_description="Choice",
typ=ToolParameter.ToolParameterType.SELECT,
required=True,
options=["a", "b"],
)
parameter.multiple = True
tool.entity.parameters = [parameter]
list(tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={"choice": ["a", "b"]}))
assert tool.last_invocation is not None
assert tool.last_invocation["tool_parameters"] == {"choice": ["a", "b"]}
with pytest.raises(ValueError, match="must be a list"):
tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={"choice": "a"})
def test_invoke_supports_list_and_generator_results():
tool = _build_tool()
tool.result = [tool.create_text_message("a"), tool.create_text_message("b")]
@ -214,6 +237,21 @@ def test_get_llm_parameters_json_schema_uses_effective_runtime_parameters():
required=False,
options=["global", "cn"],
)
regions_parameter = ToolParameter.get_simple_instance(
name="regions",
llm_description="Search regions",
typ=ToolParameter.ToolParameterType.SELECT,
required=False,
options=["global", "cn"],
)
regions_parameter.multiple = True
tags_parameter = ToolParameter.get_simple_instance(
name="tags",
llm_description="Search tags",
typ=ToolParameter.ToolParameterType.DYNAMIC_SELECT,
required=False,
)
tags_parameter.multiple = True
hidden_parameter = ToolParameter.get_simple_instance(
name="api_key",
llm_description="Hidden api key",
@ -241,7 +279,15 @@ def test_get_llm_parameters_json_schema_uses_effective_runtime_parameters():
"properties": {"nested": {"type": "string"}},
},
)
tool.entity.parameters = [query_parameter, region_parameter, hidden_parameter, file_parameter, payload_parameter]
tool.entity.parameters = [
query_parameter,
region_parameter,
regions_parameter,
tags_parameter,
hidden_parameter,
file_parameter,
payload_parameter,
]
query_override = ToolParameter.get_simple_instance(
name="query",
@ -262,6 +308,16 @@ def test_get_llm_parameters_json_schema_uses_effective_runtime_parameters():
"description": "Search region",
"enum": ["global", "cn"],
},
"regions": {
"type": "array",
"items": {"type": "string", "enum": ["global", "cn"]},
"description": "Search regions",
},
"tags": {
"type": "array",
"items": {"type": "string"},
"description": "Search tags",
},
"payload": {
"type": "object",
"properties": {"nested": {"type": "string"}},

View File

@ -1,5 +1,8 @@
import pytest
from pydantic import ValidationError
from core.tools.entities.common_entities import I18nObject
from core.tools.entities.tool_entities import ToolEntity, ToolIdentity, ToolInvokeMessage
from core.tools.entities.tool_entities import ToolEntity, ToolIdentity, ToolInvokeMessage, ToolParameter
def _make_identity() -> ToolIdentity:
@ -11,6 +14,65 @@ def _make_identity() -> ToolIdentity:
)
def _make_select_parameter(**updates: object) -> ToolParameter:
data = ToolParameter.get_simple_instance(
name="choice",
llm_description="Choice",
typ=ToolParameter.ToolParameterType.SELECT,
required=False,
options=["a", "b"],
).model_dump()
data.update(updates)
return ToolParameter.model_validate(data)
@pytest.mark.parametrize(
("updates", "message"),
[
({"type": ToolParameter.ToolParameterType.STRING, "multiple": True}, "multiple is only valid"),
({"multiple": True, "default": "a"}, "default must be a list"),
({"default": ["a"]}, "default must be a list"),
],
)
def test_tool_parameter_rejects_invalid_multiple_declarations(updates: dict[str, object], message: str):
with pytest.raises(ValidationError, match=message):
_make_select_parameter(**updates)
@pytest.mark.parametrize(
"parameter_type",
[ToolParameter.ToolParameterType.SELECT, ToolParameter.ToolParameterType.DYNAMIC_SELECT],
)
def test_tool_parameter_accepts_multiple_select_declarations(parameter_type: ToolParameter.ToolParameterType):
parameter = _make_select_parameter(type=parameter_type, multiple=True, default=["a"])
assert parameter.multiple is True
@pytest.mark.parametrize(
("value", "message"),
[
("a", "must be a list"),
(["a", 1], "only strings"),
(["missing"], "not in options"),
([], "not found in tool config"),
],
)
def test_multiple_select_normalization_rejects_invalid_values(value: object, message: str):
parameter = _make_select_parameter(multiple=True, required=True)
with pytest.raises(ValueError, match=message):
parameter.init_frontend_parameter(value)
def test_multiple_select_normalization_preserves_explicit_empty_list():
parameter = _make_select_parameter(multiple=True, default=["a"])
assert parameter.init_frontend_parameter(None) == ["a"]
assert parameter.init_frontend_parameter([]) == []
assert parameter.init_frontend_parameter(["a", "b"]) == ["a", "b"]
def test_log_message_metadata_none_defaults_to_empty_dict():
log_message = ToolInvokeMessage.LogMessage(
id="log-1",

View File

@ -6,6 +6,8 @@ from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from sqlalchemy import Engine
from sqlalchemy.orm import sessionmaker
from core.callback_handler.workflow_tool_callback_handler import DifyWorkflowCallbackHandler
from core.plugin.impl.exc import PluginDaemonClientSideError, PluginInvokeError
@ -26,7 +28,7 @@ from tests.workflow_test_utils import build_test_graph_init_params, build_test_v
@pytest.fixture
def runtime(monkeypatch) -> DifyToolNodeRuntime:
def runtime(monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> DifyToolNodeRuntime:
module_name = "core.ops.ops_trace_manager"
if module_name not in sys.modules:
ops_stub = types.ModuleType(module_name)
@ -44,9 +46,7 @@ def runtime(monkeypatch) -> DifyToolNodeRuntime:
invoke_from="debugger",
call_depth=0,
)
session_maker = MagicMock()
session_maker.begin.return_value.__enter__.return_value = MagicMock(name="session")
session_maker.begin.return_value.__exit__.return_value = None
session_maker = sessionmaker(sqlite_engine, expire_on_commit=False)
return DifyToolNodeRuntime(init_params.run_context, session_maker=session_maker)

View File

@ -1,6 +1,7 @@
from types import SimpleNamespace
from uuid import uuid4
import pytest
from sqlalchemy.orm import Session
from core.workflow.human_input_forms import (
load_form_dispositions_by_form_id,
@ -11,19 +12,24 @@ from core.workflow.human_input_policy import (
HumanInputSurface,
disposition_for_surface,
)
from models.human_input import RecipientType
from models.human_input import HumanInputFormRecipient, RecipientType
TABLES = (HumanInputFormRecipient,)
class _FakeSession:
def __init__(self, recipients: list[SimpleNamespace]) -> None:
self._recipients = recipients
def scalars(self, _stmt):
return self._recipients
def _recipient(form_id: str, recipient_type: RecipientType, access_token: str) -> HumanInputFormRecipient:
return HumanInputFormRecipient(
form_id=form_id,
delivery_id=str(uuid4()),
recipient_type=recipient_type,
recipient_payload="{}",
access_token=access_token,
)
def _recipient(form_id: str, recipient_type: RecipientType, access_token: str | None) -> SimpleNamespace:
return SimpleNamespace(form_id=form_id, recipient_type=recipient_type, access_token=access_token)
def _persist_recipients(session: Session, recipients: list[HumanInputFormRecipient]) -> None:
session.add_all(recipients)
session.commit()
@pytest.mark.parametrize(
@ -35,60 +41,75 @@ def _recipient(form_id: str, recipient_type: RecipientType, access_token: str |
(HumanInputSurface.SERVICE_API, "web-token"),
],
)
def test_load_form_tokens_picks_token_for_surface(surface, expected_token) -> None:
session = _FakeSession(
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
def test_load_form_tokens_picks_token_for_surface(surface, expected_token, sqlite_session: Session) -> None:
_persist_recipients(
sqlite_session,
[
_recipient("form-1", RecipientType.STANDALONE_WEB_APP, "web-token"),
_recipient("form-1", RecipientType.CONSOLE, "console-token"),
_recipient("form-1", RecipientType.BACKSTAGE, "backstage-token"),
]
_recipient("form-2", RecipientType.BACKSTAGE, "decoy-token"),
],
)
assert load_form_tokens_by_form_id(["form-1"], session=session, surface=surface) == {"form-1": expected_token}
assert load_form_tokens_by_form_id(["form-1"], session=sqlite_session, surface=surface) == {
"form-1": expected_token
}
def test_load_form_tokens_drops_forms_without_actionable_token() -> None:
session = _FakeSession(
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
def test_load_form_tokens_drops_forms_without_actionable_token(sqlite_session: Session) -> None:
_persist_recipients(
sqlite_session,
[
_recipient("form-1", RecipientType.EMAIL_MEMBER, "email-token"),
_recipient("form-1", RecipientType.CONSOLE, None),
]
_recipient("form-1", RecipientType.CONSOLE, ""),
],
)
assert load_form_tokens_by_form_id(["form-1"], session=session) == {}
assert load_form_tokens_by_form_id(["form-1"], session=sqlite_session) == {}
def test_load_form_tokens_service_api_surface_uses_web_token() -> None:
session = _FakeSession(
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
def test_load_form_tokens_service_api_surface_uses_web_token(sqlite_session: Session) -> None:
_persist_recipients(
sqlite_session,
[
_recipient("form-1", RecipientType.STANDALONE_WEB_APP, "web-token"),
_recipient("form-1", RecipientType.CONSOLE, "console-token"),
_recipient("form-1", RecipientType.BACKSTAGE, "backstage-token"),
]
],
)
assert load_form_tokens_by_form_id(["form-1"], session=session, surface=HumanInputSurface.SERVICE_API) == {
assert load_form_tokens_by_form_id(["form-1"], session=sqlite_session, surface=HumanInputSurface.SERVICE_API) == {
"form-1": "web-token"
}
def test_load_dispositions_openapi_webapp_form_is_resumable() -> None:
session = _FakeSession(
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
def test_load_dispositions_openapi_webapp_form_is_resumable(sqlite_session: Session) -> None:
_persist_recipients(
sqlite_session,
[
_recipient("form-1", RecipientType.STANDALONE_WEB_APP, "web-token"),
_recipient("form-1", RecipientType.BACKSTAGE, "backstage-token"),
]
],
)
assert load_form_dispositions_by_form_id(["form-1"], session=session, surface=HumanInputSurface.OPENAPI) == {
assert load_form_dispositions_by_form_id(["form-1"], session=sqlite_session, surface=HumanInputSurface.OPENAPI) == {
"form-1": FormDisposition(form_token="web-token", approval_channels=["console"])
}
def test_load_dispositions_openapi_backstage_only_form_yields_channels_not_token() -> None:
session = _FakeSession([_recipient("form-1", RecipientType.BACKSTAGE, "backstage-token")])
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
def test_load_dispositions_openapi_backstage_only_form_yields_channels_not_token(sqlite_session: Session) -> None:
_persist_recipients(
sqlite_session,
[_recipient("form-1", RecipientType.BACKSTAGE, "backstage-token")],
)
assert load_form_dispositions_by_form_id(["form-1"], session=session, surface=HumanInputSurface.OPENAPI) == {
assert load_form_dispositions_by_form_id(["form-1"], session=sqlite_session, surface=HumanInputSurface.OPENAPI) == {
"form-1": FormDisposition(form_token=None, approval_channels=["console"])
}

View File

@ -1,8 +1,12 @@
from collections.abc import Iterator
from datetime import UTC, datetime
from types import SimpleNamespace
from unittest.mock import MagicMock, Mock, sentinel
from uuid import uuid4
import pytest
from sqlalchemy import Engine, event
from sqlalchemy.orm import Session, sessionmaker
from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext, InvokeFrom, UserFrom
from core.app.file_access import FileAccessScope, bind_file_access_scope, grant_retriever_segment_access
@ -44,9 +48,62 @@ from graphon.model_runtime.model_providers.base.large_language_model import Larg
from graphon.nodes.llm.runtime_protocols import LLMPollingCapableProtocol
from graphon.nodes.tool.entities import ToolNodeData, ToolProviderType
from graphon.variables.segments import ArrayFileSegment, FileSegment
from models.base import TypeBase
from models.dataset import SegmentAttachmentBinding
from models.enums import CreatorUserRole
from models.model import StorageType, UploadFile
from tests.workflow_test_utils import build_test_run_context
@pytest.fixture
def attachment_session(sqlite_engine: Engine) -> Iterator[Session]:
"""Provide real attachment and upload-file persistence to node runtime tests."""
TypeBase.metadata.create_all(sqlite_engine, tables=[SegmentAttachmentBinding.__table__, UploadFile.__table__])
with Session(sqlite_engine, expire_on_commit=False) as session:
yield session
def _persist_attachment(
session: Session,
*,
segment_id: str,
upload_file_id: str,
upload_file_tenant_id: str = "tenant-id",
) -> UploadFile:
"""Persist an attachment binding for the test tenant and its referenced upload file."""
upload_file = UploadFile(
tenant_id=upload_file_tenant_id,
storage_type=StorageType.LOCAL,
key="storage-key",
name="diagram.png",
size=128,
extension="png",
mime_type="image/png",
created_by_role=CreatorUserRole.ACCOUNT,
created_by="user-id",
created_at=datetime.now(UTC).replace(tzinfo=None),
used=False,
source_url="https://example.com/diagram.png",
)
upload_file.id = upload_file_id
session.add_all(
[
upload_file,
SegmentAttachmentBinding(
tenant_id="tenant-id",
dataset_id="dataset-id",
document_id="document-id",
segment_id=segment_id,
attachment_id=upload_file_id,
),
]
)
session.commit()
return upload_file
def _build_model_schema(*, features: list[ModelFeature] | None = None) -> AIModelEntity:
return AIModelEntity(
model="gpt-4o-mini",
@ -348,29 +405,12 @@ def test_dify_prompt_message_serializer_delegates(monkeypatch: pytest.MonkeyPatc
)
def test_dify_retriever_attachment_loader_builds_graph_files(monkeypatch: pytest.MonkeyPatch) -> None:
upload_file = SimpleNamespace(
id="upload-file-id",
name="diagram.png",
extension="png",
mime_type="image/png",
source_url="https://example.com/diagram.png",
key="storage-key",
size=128,
)
session = MagicMock()
session.execute.return_value.all.return_value = [(None, upload_file)]
class _SessionContext:
def __enter__(self):
return session
def __exit__(self, exc_type, exc, tb):
return False
def test_dify_retriever_attachment_loader_builds_graph_files(
monkeypatch: pytest.MonkeyPatch, attachment_session: Session
) -> None:
_persist_attachment(attachment_session, segment_id="segment-id", upload_file_id="upload-file-id")
build_from_mapping = MagicMock(return_value=sentinel.file)
monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=object()))
monkeypatch.setattr(node_runtime, "Session", MagicMock(return_value=_SessionContext()))
monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=attachment_session.get_bind()))
loader = DifyRetrieverAttachmentLoader(
file_reference_factory=SimpleNamespace(build_from_mapping=build_from_mapping)
)
@ -388,39 +428,18 @@ def test_dify_retriever_attachment_loader_builds_graph_files(monkeypatch: pytest
def test_dify_retriever_attachment_loader_grants_upload_files_for_allowed_segment(
monkeypatch: pytest.MonkeyPatch,
attachment_session: Session,
) -> None:
from factories.file_factory import builders as file_builders
upload_file_id = str(uuid4())
segment_id = str(uuid4())
upload_file = SimpleNamespace(
id=upload_file_id,
tenant_id="tenant-id",
name="diagram.png",
extension="png",
mime_type="image/png",
source_url="https://example.com/diagram.png",
key="storage-key",
size=128,
)
attachment_session = MagicMock()
attachment_session.execute.return_value.all.return_value = [(None, upload_file)]
class _AttachmentSessionContext:
def __enter__(self):
return attachment_session
def __exit__(self, exc_type, exc, tb):
return False
upload_session = MagicMock()
upload_session.__enter__.return_value = upload_session
upload_session.__exit__.return_value = False
upload_session.scalar.return_value = upload_file
monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=object()))
monkeypatch.setattr(node_runtime, "Session", MagicMock(return_value=_AttachmentSessionContext()))
monkeypatch.setattr(file_builders, "session_factory", SimpleNamespace(create_session=lambda: upload_session))
_persist_attachment(attachment_session, segment_id=segment_id, upload_file_id=upload_file_id)
engine = attachment_session.get_bind()
assert engine is not None
monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=engine))
session_maker = sessionmaker(engine, expire_on_commit=False)
monkeypatch.setattr(file_builders.session_factory, "create_session", session_maker)
loader = DifyRetrieverAttachmentLoader(file_reference_factory=DifyFileReferenceFactory(_build_run_context()))
scope = FileAccessScope(
@ -435,18 +454,57 @@ def test_dify_retriever_attachment_loader_grants_upload_files_for_allowed_segmen
files = loader.load(segment_id=segment_id)
assert files[0].related_id == upload_file_id
stmt = upload_session.scalar.call_args.args[0]
whereclause = str(stmt.whereclause)
assert "upload_files.tenant_id" in whereclause
assert "upload_files.id IN" in whereclause
assert files[0].filename == "diagram.png"
def test_dify_retriever_attachment_loader_rejects_granted_upload_file_from_another_tenant(
monkeypatch: pytest.MonkeyPatch,
attachment_session: Session,
) -> None:
from factories.file_factory import builders as file_builders
upload_file_id = str(uuid4())
segment_id = str(uuid4())
_persist_attachment(
attachment_session,
segment_id=segment_id,
upload_file_id=upload_file_id,
upload_file_tenant_id="other-tenant-id",
)
engine = attachment_session.get_bind()
assert engine is not None
monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=engine))
monkeypatch.setattr(file_builders.session_factory, "create_session", sessionmaker(engine, expire_on_commit=False))
loader = DifyRetrieverAttachmentLoader(file_reference_factory=DifyFileReferenceFactory(_build_run_context()))
scope = FileAccessScope(
tenant_id="tenant-id",
user_id="end-user-id",
user_from=UserFrom.END_USER,
invoke_from=InvokeFrom.WEB_APP,
)
with bind_file_access_scope(scope):
grant_retriever_segment_access([segment_id])
with pytest.raises(ValueError, match="Invalid upload file"):
loader.load(segment_id=segment_id)
def test_dify_retriever_attachment_loader_skips_ungranted_segment_for_end_user(
monkeypatch: pytest.MonkeyPatch,
attachment_session: Session,
) -> None:
build_from_mapping = MagicMock()
session_factory = MagicMock()
monkeypatch.setattr(node_runtime, "Session", session_factory)
engine = attachment_session.get_bind()
assert engine is not None
monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=engine))
statement_count = 0
def count_statements(*_args, **_kwargs) -> None:
nonlocal statement_count
statement_count += 1
event.listen(engine, "before_cursor_execute", count_statements)
loader = DifyRetrieverAttachmentLoader(
file_reference_factory=SimpleNamespace(build_from_mapping=build_from_mapping)
)
@ -460,19 +518,31 @@ def test_dify_retriever_attachment_loader_skips_ungranted_segment_for_end_user(
with bind_file_access_scope(scope):
files = loader.load(segment_id=str(uuid4()))
assert files == []
session_factory.assert_not_called()
build_from_mapping.assert_not_called()
try:
assert files == []
assert statement_count == 0
build_from_mapping.assert_not_called()
finally:
event.remove(engine, "before_cursor_execute", count_statements)
def test_dify_retriever_attachment_loader_skips_segment_rejected_by_checker(
monkeypatch: pytest.MonkeyPatch,
attachment_session: Session,
) -> None:
segment_id = str(uuid4())
build_from_mapping = MagicMock()
session_factory = MagicMock()
segment_access_checker = MagicMock(return_value=False)
monkeypatch.setattr(node_runtime, "Session", session_factory)
engine = attachment_session.get_bind()
assert engine is not None
monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=engine))
statement_count = 0
def count_statements(*_args, **_kwargs) -> None:
nonlocal statement_count
statement_count += 1
event.listen(engine, "before_cursor_execute", count_statements)
loader = DifyRetrieverAttachmentLoader(
file_reference_factory=SimpleNamespace(build_from_mapping=build_from_mapping),
segment_access_checker=segment_access_checker,
@ -488,10 +558,13 @@ def test_dify_retriever_attachment_loader_skips_segment_rejected_by_checker(
grant_retriever_segment_access([segment_id])
files = loader.load(segment_id=segment_id)
assert files == []
segment_access_checker.assert_called_once_with(segment_id)
session_factory.assert_not_called()
build_from_mapping.assert_not_called()
try:
assert files == []
segment_access_checker.assert_called_once_with(segment_id)
assert statement_count == 0
build_from_mapping.assert_not_called()
finally:
event.remove(engine, "before_cursor_execute", count_statements)
def test_dify_tool_file_manager_resolves_conversation_id_for_tool_files(monkeypatch: pytest.MonkeyPatch) -> None:

View File

@ -2,6 +2,7 @@
import json
import os
import re
import subprocess
import sys
from pathlib import Path
@ -10,8 +11,12 @@ from typing import cast
import pytest
from dev import generate_knowledge_fs_contract as contract_validator
from dev.generate_knowledge_fs_contract import ContractDeclaration, validate_declarations
from services.knowledge_fs_proxy import KNOWLEDGE_FS_CONSOLE_OPERATIONS, KnowledgeFSOperation
from dev.generate_knowledge_fs_contract import (
ContractDeclaration,
filter_openapi_document,
validate_declarations,
)
from services.knowledge_fs_operations import KNOWLEDGE_FS_CONSOLE_OPERATIONS, KnowledgeFSOperation
def test_contract_cli_updates_checks_and_detects_openapi_drift(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
@ -55,7 +60,8 @@ def test_contract_cli_updates_checks_and_detects_openapi_drift(tmp_path: Path, m
)
)
monkeypatch.setattr(contract_validator, "LOCK_PATH", lock_path)
monkeypatch.setenv("PATH", f"{executable_directory}{os.pathsep}{os.environ['PATH']}")
current_path = os.environ.get("PATH", os.defpath)
monkeypatch.setenv("PATH", f"{executable_directory}{os.pathsep}{current_path}")
monkeypatch.setattr(
sys,
@ -106,7 +112,7 @@ def test_contract_script_loads_runtime_registry_outside_api_directory(tmp_path:
text=True,
)
assert result.stdout.strip() == "2"
assert result.stdout.strip() == str(len(KNOWLEDGE_FS_CONSOLE_OPERATIONS))
def test_validate_declarations_accepts_matching_contract() -> None:
@ -131,43 +137,157 @@ def test_validate_declarations_accepts_matching_contract() -> None:
)
def test_console_operation_registry_matches_contract() -> None:
def test_filter_openapi_document_keeps_only_declared_operations_and_referenced_schemas() -> None:
list_route = operation("knowledge-spaces:read", "listKnowledgeSpaces")
create_route = operation("knowledge-spaces:write", "createKnowledgeSpace")
for route in (list_route, create_route):
route["parameters"] = [{"in": "header", "name": "X-Trace-Id"}]
route["responses"] = {
"200": {
"content": {"application/json": {}},
"headers": {"X-Trace-Id": {}},
}
}
validate_declarations(
{
"paths": {
"/knowledge-spaces": {
"get": list_route,
"post": create_route,
list_route["responses"] = {
"200": {
"content": {
"application/json": {
"schema": {"$ref": "#/components/schemas/KnowledgeSpaceList"},
}
}
}
}
document = {
"openapi": "3.1.0",
"paths": {
"/knowledge-spaces": {
"get": list_route,
"post": operation("knowledge-spaces:write", "createKnowledgeSpace"),
},
"/health": {"get": operation(None, "getHealth", security=[])},
},
"components": {
"schemas": {
"KnowledgeSpaceList": {
"type": "object",
"properties": {
"items": {
"type": "array",
"items": {"$ref": "#/components/schemas/KnowledgeSpace"},
}
},
},
"KnowledgeSpace": {"type": "object"},
"Unused": {"type": "object"},
},
"securitySchemes": {"bearerAuth": {"type": "http", "scheme": "bearer"}},
},
}
filtered = filter_openapi_document(document, (declaration(),))
assert set(filtered["paths"]) == {"/knowledge-spaces"}
assert set(filtered["paths"]["/knowledge-spaces"]) == {"get"}
assert set(filtered["components"]["schemas"]) == {
"ConsoleProxyError",
"KnowledgeSpaceList",
"KnowledgeSpace",
}
assert filtered["components"]["securitySchemes"] == document["components"]["securitySchemes"]
def test_filter_openapi_document_keeps_sse_for_streaming_orpc_contracts() -> None:
json_declaration = declaration()
stream_declaration = declaration(
operation_id="streamTask",
path="tasks/{id}/events",
response_kind="stream",
response_media_types=("text/event-stream",),
)
json_operation = operation("knowledge-spaces:read", "listKnowledgeSpaces")
stream_operation = operation(
"knowledge-spaces:read",
"streamTask",
responses={"200": {"content": {"text/event-stream": {"schema": {"$ref": "#/components/schemas/TaskEvent"}}}}},
)
filtered = filter_openapi_document(
{
"paths": {
"/knowledge-spaces": {"get": json_operation},
"/tasks/{id}/events": {"get": stream_operation},
},
"components": {"schemas": {"TaskEvent": {"type": "object"}}},
},
(json_declaration, stream_declaration),
)
assert set(filtered["paths"]) == {"/knowledge-spaces", "/tasks/{id}/events"}
assert filtered["components"]["schemas"]["TaskEvent"] == {"type": "object"}
def test_filter_openapi_document_rewrites_proxy_error_responses() -> None:
route = operation("knowledge-spaces:read", "listKnowledgeSpaces")
route["responses"] = {
"200": {"content": {"application/json": {}}},
"401": {"content": {"application/json": {"schema": {"$ref": "#/components/schemas/ErrorResponse"}}}},
"403": {"content": {"application/json": {"schema": {"$ref": "#/components/schemas/ErrorResponse"}}}},
}
document = {
"paths": {"/knowledge-spaces": {"get": route}},
"components": {"schemas": {"ErrorResponse": {"type": "object"}}},
}
filtered = filter_openapi_document(
document,
(declaration(error_status_map=((401, 502), (403, 403))),),
)
responses = filtered["paths"]["/knowledge-spaces"]["get"]["responses"]
assert "401" not in responses
assert responses["403"]["content"]["application/json"]["schema"] == {
"$ref": "#/components/schemas/ConsoleProxyError"
}
assert responses["502"]["content"]["application/json"]["schema"] == {
"$ref": "#/components/schemas/ConsoleProxyError"
}
assert filtered["components"]["schemas"]["ConsoleProxyError"]["required"] == ["code", "message", "status"]
def test_console_operation_registry_matches_contract() -> None:
validate_declarations(
console_registry_document(),
tuple(_contract_declaration(operation) for operation in KNOWLEDGE_FS_CONSOLE_OPERATIONS),
)
def test_generated_contract_metadata_matches_current_pin_and_registry() -> None:
metadata = (
contract_validator.WORKSPACE_ROOT / "packages/contracts/generated/knowledge-fs/metadata.gen.ts"
).read_text()
lock = json.loads(contract_validator.LOCK_PATH.read_text())
assert _metadata_string(metadata, "knowledgeFsSourceOpenapiSha256") == lock["openapiSha256"]
assert _metadata_string(
metadata,
"knowledgeFsConsoleDeclarationsSha256",
) == contract_validator.contract_declarations_sha256(contract_validator.console_contract_declarations())
def _metadata_string(source: str, export_name: str) -> str:
match = re.search(rf"export const {export_name}\s*=\s*'([0-9a-f]{{64}})'", source)
assert match is not None, f"missing generated metadata export: {export_name}"
return match.group(1)
def console_registry_document() -> dict[str, object]:
list_route = operation("knowledge-spaces:read", "listKnowledgeSpaces")
create_route = operation("knowledge-spaces:write", "createKnowledgeSpace")
for route in (list_route, create_route):
route["parameters"] = [{"in": "header", "name": "X-Trace-Id"}]
route["responses"] = {
"200": {
"content": {"application/json": {}},
"headers": {"X-Trace-Id": {}},
}
}
return {"paths": {"/knowledge-spaces": {"get": list_route, "post": create_route}}}
paths: dict[str, dict[str, object]] = {}
for console_operation in KNOWLEDGE_FS_CONSOLE_OPERATIONS:
route = operation(
console_operation.required_scope,
console_operation.operation_id,
parameters=[{"in": "header", "name": name} for name in console_operation.request_headers],
responses={
"200": {
"content": {media_type: {} for media_type in console_operation.response_media_types},
"headers": {name: {} for name in console_operation.response_headers},
}
},
)
route["x-knowledge-fs-max-response-bytes"] = console_operation.max_response_bytes
paths.setdefault(f"/{console_operation.path}", {})[console_operation.method.lower()] = route
return {"paths": paths}
@pytest.mark.parametrize(
@ -326,6 +446,7 @@ def declaration(**overrides: object) -> ContractDeclaration:
"request_headers": (),
"response_headers": (),
"response_media_types": ("application/json",),
"error_status_map": ((401, 502), (403, 403)),
}
value.update(overrides)
return cast(ContractDeclaration, value)
@ -342,6 +463,7 @@ def _contract_declaration(operation: KnowledgeFSOperation) -> ContractDeclaratio
"request_headers": operation.request_headers,
"response_headers": operation.response_headers,
"response_media_types": operation.response_media_types,
"error_status_map": operation.error_status_map,
}

View File

@ -0,0 +1,27 @@
from unittest.mock import MagicMock
from extensions.storage.aws_s3_storage import AwsS3Storage
def test_generate_presigned_url() -> None:
storage = AwsS3Storage.__new__(AwsS3Storage)
storage.bucket_name = "test-bucket"
storage.client = MagicMock()
storage.client.generate_presigned_url.return_value = "https://s3.example.com/icon.png?signature=test"
result = storage.generate_presigned_url(
"upload_files/tenant/icon.png",
expires_in=300,
content_type="image/png",
)
assert result == "https://s3.example.com/icon.png?signature=test"
storage.client.generate_presigned_url.assert_called_once_with(
"get_object",
Params={
"Bucket": "test-bucket",
"Key": "upload_files/tenant/icon.png",
"ResponseContentType": "image/png",
},
ExpiresIn=300,
)

View File

@ -4,6 +4,7 @@ from werkzeug.exceptions import BadRequest, Unauthorized
from constants import COOKIE_NAME_ACCESS_TOKEN, COOKIE_NAME_CSRF_TOKEN, COOKIE_NAME_REFRESH_TOKEN
from core.errors.error import AppInvokeQuotaExceededError
from core.plugin.impl.exc import PluginRuntimeError
from libs.exception import BaseHTTPException
from libs.external_api import ExternalApi
from libs.rate_limit import _BearerRateLimited
@ -39,6 +40,14 @@ def _create_api_app():
def get(self):
raise RuntimeError("oops")
@api.route("/plugin-runtime-error")
class PluginRuntime(Resource):
def get(self):
raise PluginRuntimeError(
"Plugin runtime request failed: Runtime.ExitError: Runtime exited with error: exit status 1",
lambda_request_id="lambda-request-id",
)
# Note: We avoid altering default_mediatype to keep normal error paths
# Special 400 message rewrite
@ -107,6 +116,24 @@ def test_external_api_json_message_and_bad_request_rewrite():
assert res.get_json()["message"] == "Invalid JSON payload received or JSON payload is empty."
def test_external_api_plugin_runtime_error(mocker):
mocker.patch("libs.external_api.get_request_id", return_value="api-request-id")
app = _create_api_app()
res = app.test_client().get("/api/plugin-runtime-error")
assert res.status_code == 502
assert res.get_json() == {
"code": "plugin_runtime_error",
"message": "Plugin runtime request failed: Runtime.ExitError: Runtime exited with error: exit status 1",
"details": {
"request_id": "api-request-id",
"lambda_request_id": "lambda-request-id",
},
"status": 502,
}
def test_external_api_param_mapping_and_quota():
app = _create_api_app()
client = app.test_client()

View File

@ -2,7 +2,7 @@ from datetime import datetime
import pytest
from libs.helper import OptionalTimestampField, escape_like_pattern, extract_tenant_id
from libs.helper import OptionalTimestampField, email, escape_like_pattern, extract_tenant_id
from models.account import Account
from models.model import EndUser
@ -126,3 +126,30 @@ class TestEscapeLikePattern:
result = escape_like_pattern("test\\%_value")
# Should be: test\\\%\_value
assert result == "test\\\\\\%\\_value"
class TestEmailValidator:
"""Tests for the email() validator — regression for #39234."""
def test_valid_email_accepted(self):
assert email("user@example.com") == "user@example.com"
def test_trailing_newline_rejected(self):
with pytest.raises(ValueError, match="not a valid email"):
email("user@example.com\n")
def test_trailing_carriage_return_newline_rejected(self):
with pytest.raises(ValueError, match="not a valid email"):
email("user@example.com\r\n")
def test_multiple_newlines_rejected(self):
with pytest.raises(ValueError, match="not a valid email"):
email("user@example.com\n\n")
def test_empty_string_rejected(self):
with pytest.raises(ValueError, match="not a valid email"):
email("")
def test_invalid_email_rejected(self):
with pytest.raises(ValueError, match="not a valid email"):
email("not-an-email")

View File

@ -139,6 +139,23 @@ def test_declared_output_child_validates_shape_and_defaults() -> None:
)
def test_declared_output_child_schema_matches_nullable_serialization() -> None:
config = DeclaredOutputConfig(
name="response",
type=DeclaredOutputType.OBJECT,
children=[DeclaredOutputChildConfig(name="text", type=DeclaredOutputType.STRING)],
)
child = config.model_dump(mode="json")["children"][0]
assert child["file"] is None
assert child["array_item"] is None
children_schema = DeclaredOutputConfig.model_json_schema(mode="serialization")["properties"]["children"]
child_properties = children_schema["items"]["properties"]
assert {"type": "null"} in child_properties["file"]["anyOf"]
assert {"type": "null"} in child_properties["array_item"]["anyOf"]
def test_declared_output_validates_shape_and_defaults() -> None:
file_output = DeclaredOutputConfig(name="report", type=DeclaredOutputType.FILE)
assert file_output.file is not None

View File

@ -1271,11 +1271,32 @@ def test_node_job_only_updates_inline_agent_soul(monkeypatch: pytest.MonkeyPatch
tenant_id="tenant-1",
agent_id="inline-agent-1",
version=2,
config_snapshot=AgentSoulConfig.model_validate(
{
"model": {
"plugin_id": "langgenius/openai/openai",
"model_provider": "openai",
"model": "gpt-4o",
},
"prompt": {"system_prompt": "new"},
}
),
)
normal_draft = AgentConfigDraft(
id="draft-1",
tenant_id="tenant-1",
agent_id="inline-agent-1",
draft_type=AgentConfigDraftType.DRAFT,
account_id=None,
draft_owner_key="",
base_snapshot_id="inline-version-1",
config_snapshot=AgentSoulConfig.model_validate({"prompt": {"system_prompt": "old"}}),
)
monkeypatch.setattr(AgentComposerService, "_require_version", lambda **kwargs: current_snapshot)
monkeypatch.setattr(AgentComposerService, "_update_current_version", lambda **kwargs: next_snapshot)
monkeypatch.setattr(AgentComposerService, "_require_agent", lambda **kwargs: inline_agent)
monkeypatch.setattr(AgentComposerService, "_get_agent_draft", lambda **kwargs: normal_draft)
binding = WorkflowAgentNodeBinding(
tenant_id="tenant-1",
@ -1320,6 +1341,95 @@ def test_node_job_only_updates_inline_agent_soul(monkeypatch: pytest.MonkeyPatch
assert inline_agent.active_config_snapshot_id == "inline-version-2"
assert inline_agent.active_config_has_model is True
assert inline_agent.updated_by == "account-1"
assert normal_draft.id == "draft-1"
assert normal_draft.base_snapshot_id == "inline-version-2"
assert normal_draft.config_snapshot_dict == next_snapshot.config_snapshot_dict
assert normal_draft.updated_by == "account-1"
def test_get_or_create_normal_agent_draft_rebases_stale_workflow_only_draft():
agent = Agent(
id="inline-agent-1",
tenant_id="tenant-1",
name="Inline",
description="",
agent_kind=AgentKind.DIFY_AGENT,
scope=AgentScope.WORKFLOW_ONLY,
source=AgentSource.WORKFLOW,
status=AgentStatus.ACTIVE,
active_config_snapshot_id="inline-version-2",
created_by="account-1",
updated_by="account-2",
)
draft = AgentConfigDraft(
id="draft-1",
tenant_id="tenant-1",
agent_id=agent.id,
draft_type=AgentConfigDraftType.DRAFT,
account_id=None,
draft_owner_key="",
base_snapshot_id="inline-version-1",
config_snapshot=AgentSoulConfig.model_validate({"prompt": {"system_prompt": "old"}}),
)
active_snapshot = AgentConfigSnapshot(
id="inline-version-2",
tenant_id="tenant-1",
agent_id=agent.id,
version=2,
config_snapshot=AgentSoulConfig.model_validate({"prompt": {"system_prompt": "new"}}),
)
session = FakeSession(scalar=[draft, active_snapshot])
resolved = AgentComposerService.get_or_create_normal_agent_draft(
session=session,
tenant_id="tenant-1",
agent=agent,
created_by="account-2",
)
assert resolved is draft
assert resolved.id == "draft-1"
assert resolved.base_snapshot_id == "inline-version-2"
assert resolved.config_snapshot_dict == active_snapshot.config_snapshot_dict
assert resolved.updated_by == "account-2"
assert session.flushes == 1
def test_get_or_create_normal_agent_draft_keeps_roster_draft_edits():
agent = Agent(
id="roster-agent-1",
tenant_id="tenant-1",
name="Roster",
description="",
agent_kind=AgentKind.DIFY_AGENT,
scope=AgentScope.ROSTER,
source=AgentSource.AGENT_APP,
status=AgentStatus.ACTIVE,
active_config_snapshot_id="version-2",
)
draft = AgentConfigDraft(
id="draft-1",
tenant_id="tenant-1",
agent_id=agent.id,
draft_type=AgentConfigDraftType.DRAFT,
account_id=None,
draft_owner_key="",
base_snapshot_id="version-1",
config_snapshot=AgentSoulConfig.model_validate({"prompt": {"system_prompt": "local edit"}}),
)
session = FakeSession(scalar=[draft])
resolved = AgentComposerService.get_or_create_normal_agent_draft(
session=session,
tenant_id="tenant-1",
agent=agent,
created_by="account-1",
)
assert resolved is draft
assert resolved.base_snapshot_id == "version-1"
assert resolved.config_snapshot_dict["prompt"]["system_prompt"] == "local edit"
assert session.flushes == 0
def test_node_job_only_switches_roster_binding_to_inline_agent(monkeypatch: pytest.MonkeyPatch):
@ -2624,7 +2734,7 @@ def test_roster_create_detail_and_lookup_helpers(monkeypatch: pytest.MonkeyPatch
monkeypatch.setattr(
AgentRosterService,
"_get_or_create_agent_app_debug_conversation",
lambda self, *, agent, account_id: "debug-conversation-1",
lambda self, *, agent, account_id, draft_type: "debug-conversation-1",
)
payload = roster_service.RosterAgentCreatePayload(
name="Analyst",
@ -2730,6 +2840,7 @@ def test_agent_app_debug_conversation_create_reuse_and_recreate():
assert created_mapping.tenant_id == "tenant-1"
assert created_mapping.agent_id == "agent-1"
assert created_mapping.account_id == "account-1"
assert created_mapping.draft_type == AgentConfigDraftType.DEBUG_BUILD
assert create_session.commits == 1
existing_mapping = AgentDebugConversation(
@ -2737,6 +2848,7 @@ def test_agent_app_debug_conversation_create_reuse_and_recreate():
agent_id="agent-1",
app_id="app-1",
account_id="account-1",
draft_type=AgentConfigDraftType.DEBUG_BUILD,
conversation_id="existing-conversation",
)
reuse_session = FakeSession(scalar=[agent, existing_mapping, "existing-conversation"])
@ -2754,6 +2866,7 @@ def test_agent_app_debug_conversation_create_reuse_and_recreate():
agent_id="agent-1",
app_id="app-1",
account_id="account-1",
draft_type=AgentConfigDraftType.DEBUG_BUILD,
conversation_id="deleted-conversation",
)
recreate_session = FakeSession(scalar=[agent, stale_mapping, None])
@ -2768,6 +2881,42 @@ def test_agent_app_debug_conversation_create_reuse_and_recreate():
assert recreate_session.commits == 1
def test_agent_app_debug_conversations_are_isolated_by_draft_type():
agent = Agent(
id="agent-1",
tenant_id="tenant-1",
app_id="app-1",
name="Analyst",
description="",
agent_kind=AgentKind.DIFY_AGENT,
scope=AgentScope.ROSTER,
source=AgentSource.AGENT_APP,
status=AgentStatus.ACTIVE,
)
session = FakeSession(scalar=[agent, None, agent, None])
service = AgentRosterService(session)
build_conversation_id = service.get_or_create_agent_app_debug_conversation_id(
tenant_id="tenant-1",
agent_id="agent-1",
account_id="account-1",
draft_type=AgentConfigDraftType.DEBUG_BUILD,
)
preview_conversation_id = service.get_or_create_agent_app_debug_conversation_id(
tenant_id="tenant-1",
agent_id="agent-1",
account_id="account-1",
draft_type=AgentConfigDraftType.DRAFT,
)
mappings = [value for value in session.added if isinstance(value, AgentDebugConversation)]
assert build_conversation_id != preview_conversation_id
assert {mapping.draft_type for mapping in mappings} == {
AgentConfigDraftType.DRAFT,
AgentConfigDraftType.DEBUG_BUILD,
}
def test_agent_app_debug_conversation_message_count():
session = FakeSession(scalar=[3])
@ -2794,6 +2943,7 @@ def test_agent_app_debug_conversation_requires_app_binding():
AgentRosterService(FakeSession())._get_or_create_agent_app_debug_conversation(
agent=agent,
account_id="account-1",
draft_type=AgentConfigDraftType.DEBUG_BUILD,
)
@ -2840,7 +2990,9 @@ def test_load_or_create_agent_app_debug_conversations_supports_runtime_backed_ag
assert result["agent-1"]
assert result["agent-3"]
assert fake_session.commits == 1
assert len([value for value in fake_session.added if isinstance(value, AgentDebugConversation)]) == 2
mappings = [value for value in fake_session.added if isinstance(value, AgentDebugConversation)]
assert len(mappings) == 2
assert all(mapping.draft_type == AgentConfigDraftType.DEBUG_BUILD for mapping in mappings)
def test_agent_app_visible_versions_exclude_draft_saves():
@ -3276,6 +3428,7 @@ class TestAgentAppBackingAgent:
assert mappings[0].agent_id == "agent-1"
assert mappings[0].app_id == "app-1"
assert mappings[0].account_id == "account-1"
assert mappings[0].draft_type == AgentConfigDraftType.DEBUG_BUILD
assert mappings[0].conversation_id == conversation_id
assert session.deleted == []
assert session.commits == 1
@ -3361,8 +3514,10 @@ class TestAgentAppBackingAgent:
payload = cleanup_delay.call_args.args[0]
assert payload["metadata"]["conversation_id"] == "old-conversation"
assert payload["metadata"]["agent_id"] == "agent-9"
assert payload["metadata"]["draft_type"] == "debug_build"
assert (
payload["idempotency_key"] == "tenant-1:agent-1:account-1:old-conversation:debug-session-cleanup:"
payload["idempotency_key"]
== "tenant-1:agent-1:account-1:debug_build:old-conversation:debug-session-cleanup:"
"agent-9:snap-9:run-old"
)
cleanup_store.mark_cleaned.assert_called_once_with(

View File

@ -1,7 +1,8 @@
from unittest.mock import MagicMock
import pytest
from sqlalchemy import event
from sqlalchemy.orm import Session
from models.tools import MCPToolProvider
from services.data_migration.dependency_discovery_service import DiscoveredDependency
from services.data_migration.entities import (
ConflictStrategy,
@ -12,6 +13,10 @@ from services.data_migration.entities import (
)
from services.data_migration.export_service import ExportConfigParser, MigrationExportService
_TENANT_ID = "11111111-1111-1111-1111-111111111111"
_OTHER_TENANT_ID = "22222222-2222-2222-2222-222222222222"
_USER_ID = "33333333-3333-3333-3333-333333333333"
def test_export_config_parser_accepts_new_scripted_shape():
selection = ExportConfigParser().parse(
@ -121,7 +126,8 @@ def test_secret_free_api_tool_export_uses_masking_and_omits_credentials(monkeypa
assert report_items[0].resource_type == ResourceType.API_TOOL
def test_secret_free_mcp_dependencies_are_dependency_only():
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_secret_free_mcp_dependencies_are_dependency_only(sqlite_session: Session):
service = MigrationExportService()
dependencies: list[dict] = []
mcp_tools: list[dict] = []
@ -134,9 +140,10 @@ def test_secret_free_mcp_dependencies_are_dependency_only():
exported_mcp_tools=mcp_tools,
dependencies=dependencies,
report_items=report_items,
session=MagicMock(),
session=sqlite_session,
)
assert not sqlite_session.in_transaction()
assert mcp_tools == []
assert dependencies == [
{
@ -150,17 +157,36 @@ def test_secret_free_mcp_dependencies_are_dependency_only():
assert report_items[0].name == "mcp_tool mcp-1"
def test_get_mcp_provider_does_not_compare_non_uuid_identifier_to_uuid_id():
@pytest.mark.parametrize("sqlite_session", [(MCPToolProvider,)], indirect=True)
def test_get_mcp_provider_does_not_compare_non_uuid_identifier_to_uuid_id(sqlite_session: Session):
sqlite_session.add(
MCPToolProvider(
name="Other tenant provider",
server_identifier="my-test-mcp",
server_url="https://example.com/mcp",
server_url_hash="other-tenant-provider",
icon=None,
tenant_id=_OTHER_TENANT_ID,
user_id=_USER_ID,
authed=False,
tools="[]",
)
)
sqlite_session.commit()
statements = []
def capture_scalar(statement):
statements.append(str(statement))
def capture_statement(_conn, _cursor, statement, _parameters, _context, _executemany):
statements.append(statement)
session = MagicMock()
session.scalar.side_effect = capture_scalar
bind = sqlite_session.get_bind()
event.listen(bind, "before_cursor_execute", capture_statement)
with pytest.raises(MigrationDataError, match="MCP provider not found"):
MigrationExportService()._get_mcp_provider("tenant-1", "my-test-mcp", session=session)
try:
with pytest.raises(MigrationDataError, match="MCP provider not found"):
MigrationExportService()._get_mcp_provider(_TENANT_ID, "my-test-mcp", session=sqlite_session)
finally:
event.remove(bind, "before_cursor_execute", capture_statement)
assert len(statements) == 1
assert "tool_mcp_providers.id =" not in statements[0]

View File

@ -1,11 +1,9 @@
"""Unit tests for services.enterprise.rbac_service.
The enterprise RBAC client is almost pure glue: each method turns a single
``EnterpriseRequest.send_inner_rbac_request`` call into a pydantic response
model. Rather than spinning up an HTTP server we monkeypatch that helper and
assert on the arguments it received; that catches both routing regressions
(wrong method / wrong path / wrong params) and model-shape regressions in
one place.
Most enterprise RBAC methods turn a single ``EnterpriseRequest.send_inner_rbac_request``
call into a pydantic response model. Rather than spinning up an HTTP server, these tests
monkeypatch that helper and assert on the request arguments and response shape. The legacy
fallbacks use SQLite to verify their database reads and committed role updates.
"""
from __future__ import annotations
@ -15,7 +13,10 @@ from unittest.mock import MagicMock, patch
import pytest
from flask import Flask
from sqlalchemy import select
from sqlalchemy.orm import Session
from models import TenantAccountJoin
from services.enterprise import rbac_service as svc
MODULE = "services.enterprise.rbac_service"
@ -533,8 +534,9 @@ class TestWorkspaceAccess:
assert call.params == {"language": "en"}
@pytest.mark.parametrize("sqlite_session", [(TenantAccountJoin,)], indirect=True)
class TestMyPermissions:
def test_resource_snapshot_maps_defaults_and_overrides(self):
def test_resource_snapshot_maps_defaults_and_overrides(self, sqlite_session: Session):
snapshot = svc.ResourcePermissionSnapshot(
default_permission_keys=["app.acl.view_layout"],
overrides=[
@ -550,7 +552,7 @@ class TestMyPermissions:
"app-2": ["app.acl.view_layout", "app.acl.edit"],
}
def test_get_without_payload_uses_get(self, mock_send: MagicMock):
def test_get_without_payload_uses_get(self, mock_send: MagicMock, sqlite_session: Session):
mock_send.return_value = {
"workspace": {"permission_keys": ["workspace.member.manage"]},
"app": {"default_permission_keys": ["app.acl.view_layout", "app.acl.test_and_run"], "overrides": []},
@ -558,7 +560,7 @@ class TestMyPermissions:
}
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True):
out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=MagicMock())
out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=sqlite_session)
call = _call_args(mock_send)
assert call.method == "GET"
@ -609,12 +611,14 @@ class TestMyPermissions:
workspace_keys: list[str],
app_keys: list[str],
dataset_keys: list[str],
sqlite_session: Session,
):
mock_session = MagicMock()
mock_session.__enter__.return_value = mock_session
mock_session.scalar.return_value = role
sqlite_session.add(
TenantAccountJoin(tenant_id="tenant-1", account_id="acct-1", role=svc.TenantAccountRole(role))
)
sqlite_session.commit()
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False):
out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=mock_session)
out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=sqlite_session)
mock_send.assert_not_called()
assert out.workspace.permission_keys == workspace_keys
@ -648,12 +652,14 @@ class TestMyPermissions:
mock_send: MagicMock,
role: str,
expected_snippet_keys: set[str],
sqlite_session: Session,
):
mock_session = MagicMock()
mock_session.__enter__.return_value = mock_session
mock_session.scalar.return_value = role
sqlite_session.add(
TenantAccountJoin(tenant_id="tenant-1", account_id="acct-1", role=svc.TenantAccountRole(role))
)
sqlite_session.commit()
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False):
out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=mock_session)
out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=sqlite_session)
actual_snippet_keys = {
permission_key for permission_key in out.workspace.permission_keys if permission_key.startswith("snippets.")
@ -662,19 +668,16 @@ class TestMyPermissions:
mock_send.assert_not_called()
assert actual_snippet_keys == expected_snippet_keys
def test_get_returns_empty_when_role_missing_and_rbac_disabled(self, mock_send: MagicMock):
mock_session = MagicMock()
mock_session.__enter__.return_value = mock_session
mock_session.scalar.return_value = None
def test_get_returns_empty_when_role_missing_and_rbac_disabled(self, mock_send: MagicMock, sqlite_session: Session):
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False):
out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=mock_session)
out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=sqlite_session)
mock_send.assert_not_called()
assert out.workspace.permission_keys == []
assert out.app.default_permission_keys == []
assert out.dataset.default_permission_keys == []
def test_get_with_single_resource_filters(self, mock_send: MagicMock):
def test_get_with_single_resource_filters(self, mock_send: MagicMock, sqlite_session: Session):
mock_send.return_value = {
"workspace": {"permission_keys": []},
"app": {
@ -685,7 +688,7 @@ class TestMyPermissions:
}
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True):
out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", app_id="app-1", session=MagicMock())
out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", app_id="app-1", session=sqlite_session)
call = _call_args(mock_send)
assert call.method == "GET"
@ -694,8 +697,9 @@ class TestMyPermissions:
assert out.app.overrides[0].resource_id == "app-1"
@pytest.mark.parametrize("sqlite_session", [(TenantAccountJoin,)], indirect=True)
class TestMemberRoles:
def test_get(self, mock_send: MagicMock):
def test_get(self, mock_send: MagicMock, sqlite_session: Session):
mock_send.return_value = {
"account_id": "acct-2",
"roles": [
@ -707,7 +711,7 @@ class TestMemberRoles:
],
}
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True):
out = svc.RBACService.MemberRoles.get("tenant-1", "acct-1", "acct-2", session=MagicMock())
out = svc.RBACService.MemberRoles.get("tenant-1", "acct-1", "acct-2", session=sqlite_session)
call = _call_args(mock_send)
assert call.method == "GET"
assert call.endpoint == "/rbac/members/rbac-roles"
@ -715,12 +719,14 @@ class TestMemberRoles:
assert out.account_id == "acct-2"
assert out.roles[0].name == "Member"
def test_get_legacy_role_includes_permission_keys(self, mock_send: MagicMock):
session = MagicMock()
session.scalar.return_value = svc.TenantAccountRole.EDITOR
def test_get_legacy_role_includes_permission_keys(self, mock_send: MagicMock, sqlite_session: Session):
sqlite_session.add(
TenantAccountJoin(tenant_id="tenant-1", account_id="acct-2", role=svc.TenantAccountRole.EDITOR)
)
sqlite_session.commit()
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False):
out = svc.RBACService.MemberRoles.get("tenant-1", "acct-1", "acct-2", session=session)
out = svc.RBACService.MemberRoles.get("tenant-1", "acct-1", "acct-2", session=sqlite_session)
mock_send.assert_not_called()
assert out.account_id == "acct-2"
@ -738,7 +744,7 @@ class TestMemberRoles:
assert "app.acl.preview" in out.roles[0].permission_keys
assert "dataset.acl.preview" in out.roles[0].permission_keys
def test_replace(self, mock_send: MagicMock):
def test_replace(self, mock_send: MagicMock, sqlite_session: Session):
mock_send.return_value = {"account_id": "acct-2", "roles": []}
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True):
svc.RBACService.MemberRoles.replace(
@ -746,7 +752,7 @@ class TestMemberRoles:
"acct-1",
"acct-2",
role_ids=["workspace.owner", "workspace.editor"],
session=MagicMock(),
session=sqlite_session,
)
call = _call_args(mock_send)
assert call.method == "PUT"
@ -754,43 +760,59 @@ class TestMemberRoles:
assert call.params == {"account_id": "acct-2"}
assert call.json == {"role_ids": ["workspace.owner", "workspace.editor"]}
def test_replace_updates_legacy_join_role_when_rbac_disabled(self, mock_send: MagicMock):
session = MagicMock()
session.__enter__.return_value = session
target_join = SimpleNamespace(role=svc.TenantAccountRole.NORMAL, account_id="acct-2")
session.scalar.return_value = target_join
def test_replace_commits_legacy_join_role_when_rbac_disabled(self, mock_send: MagicMock, sqlite_session: Session):
target_join = TenantAccountJoin(tenant_id="tenant-1", account_id="acct-2", role=svc.TenantAccountRole.NORMAL)
sqlite_session.add(target_join)
sqlite_session.commit()
target_join_id = target_join.id
engine = sqlite_session.get_bind()
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False):
out = svc.RBACService.MemberRoles.replace(
"tenant-1", "acct-1", "acct-2", role_ids=["editor"], session=session
"tenant-1", "acct-1", "acct-2", role_ids=["editor"], session=sqlite_session
)
mock_send.assert_not_called()
session.commit.assert_called_once()
assert target_join.role == svc.TenantAccountRole.EDITOR
# Closing the writer rolls back any uncommitted update and prevents its identity map
# from satisfying the verification query.
sqlite_session.close()
with Session(engine) as verification_session:
persisted_join = verification_session.scalar(
select(TenantAccountJoin).where(TenantAccountJoin.id == target_join_id)
)
assert persisted_join is not None
assert persisted_join.role == svc.TenantAccountRole.EDITOR
assert out.account_id == "acct-2"
assert out.roles[0].id == "editor"
assert "app.acl.preview" in out.roles[0].permission_keys
def test_replace_legacy_owner_demotes_current_owner_when_rbac_disabled(self, mock_send: MagicMock):
session = MagicMock()
session.__enter__.return_value = session
target_join = SimpleNamespace(role=svc.TenantAccountRole.NORMAL, account_id="acct-2")
owner_join = SimpleNamespace(role=svc.TenantAccountRole.OWNER, account_id="acct-owner")
session.scalar.side_effect = [target_join, owner_join]
def test_replace_legacy_owner_demotes_current_owner_when_rbac_disabled(
self, mock_send: MagicMock, sqlite_session: Session
):
target_join = TenantAccountJoin(tenant_id="tenant-1", account_id="acct-2", role=svc.TenantAccountRole.NORMAL)
owner_join = TenantAccountJoin(tenant_id="tenant-1", account_id="acct-owner", role=svc.TenantAccountRole.OWNER)
sqlite_session.add_all([target_join, owner_join])
sqlite_session.commit()
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False):
out = svc.RBACService.MemberRoles.replace(
"tenant-1", "acct-1", "acct-2", role_ids=["owner"], session=session
"tenant-1", "acct-1", "acct-2", role_ids=["owner"], session=sqlite_session
)
mock_send.assert_not_called()
session.commit.assert_called_once()
assert target_join.role == svc.TenantAccountRole.OWNER
assert owner_join.role == svc.TenantAccountRole.ADMIN
persisted_joins = {
join.account_id: join.role
for join in sqlite_session.scalars(
select(TenantAccountJoin).where(TenantAccountJoin.tenant_id == "tenant-1")
)
}
assert persisted_joins == {
"acct-2": svc.TenantAccountRole.OWNER,
"acct-owner": svc.TenantAccountRole.ADMIN,
}
assert out.roles[0].id == "owner"
def test_batch_get(self, mock_send: MagicMock):
def test_batch_get(self, mock_send: MagicMock, sqlite_session: Session):
mock_send.return_value = {
"acct-2": [
{"id": "role-1", "name": "Admin", "type": "workspace"},
@ -811,8 +833,9 @@ class TestMemberRoles:
assert out[1].roles == []
@pytest.mark.parametrize("sqlite_session", [(TenantAccountJoin,)], indirect=True)
class TestResourcePermissions:
def test_app_permissions_batch_get(self, mock_send: MagicMock):
def test_app_permissions_batch_get(self, mock_send: MagicMock, sqlite_session: Session):
mock_send.return_value = {
"data": [
{"resource_id": "app-1", "permission_keys": ["app.acl.view_layout", "app.acl.edit"]},
@ -822,7 +845,7 @@ class TestResourcePermissions:
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True):
out = svc.RBACService.AppPermissions.batch_get(
"tenant-1", "acct-1", ["app-1", "app-2"], session=MagicMock()
"tenant-1", "acct-1", ["app-1", "app-2"], session=sqlite_session
)
call = _call_args(mock_send)
@ -834,13 +857,16 @@ class TestResourcePermissions:
"app-2": [],
}
def test_app_permissions_batch_get_uses_legacy_role_permissions_when_rbac_disabled(self, mock_send: MagicMock):
mock_session = MagicMock()
mock_session.__enter__.return_value = mock_session
mock_session.scalar.return_value = "editor"
def test_app_permissions_batch_get_uses_legacy_role_permissions_when_rbac_disabled(
self, mock_send: MagicMock, sqlite_session: Session
):
sqlite_session.add(
TenantAccountJoin(tenant_id="tenant-1", account_id="acct-1", role=svc.TenantAccountRole.EDITOR)
)
sqlite_session.commit()
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False):
out = svc.RBACService.AppPermissions.batch_get(
"tenant-1", "acct-1", ["app-1", "app-2"], session=mock_session
"tenant-1", "acct-1", ["app-1", "app-2"], session=sqlite_session
)
mock_send.assert_not_called()
@ -849,7 +875,7 @@ class TestResourcePermissions:
"app-2": svc._LEGACY_APP_EDITOR_KEYS,
}
def test_dataset_permissions_batch_get(self, mock_send: MagicMock):
def test_dataset_permissions_batch_get(self, mock_send: MagicMock, sqlite_session: Session):
mock_send.return_value = {
"data": [
{"resource_id": "ds-1", "permission_keys": ["dataset.acl.readonly"]},
@ -859,7 +885,7 @@ class TestResourcePermissions:
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True):
out = svc.RBACService.DatasetPermissions.batch_get(
"tenant-1", "acct-1", ["ds-1", "ds-2"], session=MagicMock()
"tenant-1", "acct-1", ["ds-1", "ds-2"], session=sqlite_session
)
call = _call_args(mock_send)
@ -871,13 +897,20 @@ class TestResourcePermissions:
"ds-2": ["dataset.acl.edit"],
}
def test_dataset_permissions_batch_get_uses_legacy_role_permissions_when_rbac_disabled(self, mock_send: MagicMock):
mock_session = MagicMock()
mock_session.__enter__.return_value = mock_session
mock_session.scalar.return_value = "dataset_operator"
def test_dataset_permissions_batch_get_uses_legacy_role_permissions_when_rbac_disabled(
self, mock_send: MagicMock, sqlite_session: Session
):
sqlite_session.add(
TenantAccountJoin(
tenant_id="tenant-1",
account_id="acct-1",
role=svc.TenantAccountRole.DATASET_OPERATOR,
)
)
sqlite_session.commit()
with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False):
out = svc.RBACService.DatasetPermissions.batch_get(
"tenant-1", "acct-1", ["ds-1", "ds-2"], session=mock_session
"tenant-1", "acct-1", ["ds-1", "ds-2"], session=sqlite_session
)
mock_send.assert_not_called()

View File

@ -6,17 +6,25 @@ which handles retrieval testing operations for datasets, including internal
dataset retrieval and external knowledge base retrieval.
"""
import json
from typing import Any
from unittest.mock import MagicMock, Mock, patch
import pytest
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.rag.models.document import Document
from core.rag.retrieval.retrieval_methods import RetrievalMethod
from models import Account
from models.dataset import Dataset
from models.dataset import Dataset, DatasetQuery
from services.hit_testing_service import HitTestingService
pytestmark = [
pytest.mark.usefixtures("sqlite_session"),
pytest.mark.parametrize("sqlite_session", [(DatasetQuery,)], indirect=True),
]
class HitTestingTestDataFactory:
"""
@ -139,17 +147,7 @@ class TestHitTestingServiceRetrieve:
various retrieval model configurations, metadata filtering, and query logging.
"""
@pytest.fixture
def mock_db_session(self):
"""
Mock database session.
Provides a mocked database session for testing database operations
like adding and committing DatasetQuery records.
"""
return MagicMock()
def test_retrieve_success_with_default_retrieval_model(self, mock_db_session):
def test_retrieve_success_with_default_retrieval_model(self, sqlite_session: Session):
"""
Test successful retrieval with default retrieval model.
@ -186,17 +184,20 @@ class TestHitTestingServiceRetrieve:
# Act
result = HitTestingService.retrieve(
dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session
dataset, query, account, retrieval_model, external_retrieval_model, session=sqlite_session
)
# Assert
assert result["query"]["content"] == query
assert len(result["records"]) == 2
mock_retrieve.assert_called_once()
mock_db_session.add.assert_called_once()
mock_db_session.commit.assert_called_once()
query_log = sqlite_session.scalar(select(DatasetQuery))
assert query_log is not None
assert query_log.dataset_id == dataset.id
assert query_log.created_by == account.id
assert json.loads(query_log.content) == [{"content_type": "text_query", "content": query}]
def test_retrieve_success_with_custom_retrieval_model(self, mock_db_session):
def test_retrieve_success_with_custom_retrieval_model(self, sqlite_session: Session):
"""
Test successful retrieval with custom retrieval model.
@ -234,7 +235,7 @@ class TestHitTestingServiceRetrieve:
# Act
result = HitTestingService.retrieve(
dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session
dataset, query, account, retrieval_model, external_retrieval_model, session=sqlite_session
)
# Assert
@ -246,7 +247,7 @@ class TestHitTestingServiceRetrieve:
assert call_kwargs["score_threshold"] == 0.7
assert call_kwargs["reranking_model"] == retrieval_model["reranking_model"]
def test_retrieve_with_metadata_filtering(self, mock_db_session):
def test_retrieve_with_metadata_filtering(self, sqlite_session: Session):
"""
Test retrieval with metadata filtering conditions.
@ -292,7 +293,7 @@ class TestHitTestingServiceRetrieve:
# Act
result = HitTestingService.retrieve(
dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session
dataset, query, account, retrieval_model, external_retrieval_model, session=sqlite_session
)
# Assert
@ -301,7 +302,7 @@ class TestHitTestingServiceRetrieve:
call_kwargs = mock_retrieve.call_args[1]
assert call_kwargs["document_ids_filter"] == ["doc-1", "doc-2"]
def test_retrieve_with_metadata_filtering_no_documents(self, mock_db_session):
def test_retrieve_with_metadata_filtering_no_documents(self, sqlite_session: Session):
"""
Test retrieval with metadata filtering that returns no documents.
@ -337,14 +338,14 @@ class TestHitTestingServiceRetrieve:
# Act
result = HitTestingService.retrieve(
dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session
dataset, query, account, retrieval_model, external_retrieval_model, session=sqlite_session
)
# Assert
assert result["query"]["content"] == query
assert result["records"] == []
def test_retrieve_with_dataset_retrieval_model(self, mock_db_session):
def test_retrieve_with_dataset_retrieval_model(self, sqlite_session: Session):
"""
Test retrieval using dataset's retrieval model when not provided.
@ -380,7 +381,7 @@ class TestHitTestingServiceRetrieve:
# Act
result = HitTestingService.retrieve(
dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session
dataset, query, account, retrieval_model, external_retrieval_model, session=sqlite_session
)
# Assert
@ -398,17 +399,7 @@ class TestHitTestingServiceExternalRetrieve:
including query escaping, response formatting, and provider validation.
"""
@pytest.fixture
def mock_db_session(self):
"""
Mock database session.
Provides a mocked database session for testing database operations
like adding and committing DatasetQuery records.
"""
return MagicMock()
def test_external_retrieve_success(self, mock_db_session):
def test_external_retrieve_success(self, sqlite_session: Session):
"""
Test successful external retrieval.
@ -443,7 +434,7 @@ class TestHitTestingServiceExternalRetrieve:
account,
external_retrieval_model,
metadata_filtering_conditions,
session=mock_db_session,
session=sqlite_session,
)
# Assert
@ -455,10 +446,13 @@ class TestHitTestingServiceExternalRetrieve:
mock_external_retrieve.assert_called_once()
# Verify query was escaped
assert mock_external_retrieve.call_args[1]["query"] == 'test query with \\"quotes\\"'
mock_db_session.add.assert_called_once()
mock_db_session.commit.assert_called_once()
query_log = sqlite_session.scalar(select(DatasetQuery))
assert query_log is not None
assert query_log.dataset_id == dataset.id
assert query_log.content == query
assert query_log.created_by == account.id
def test_external_retrieve_non_external_provider(self, mock_db_session):
def test_external_retrieve_non_external_provider(self, sqlite_session: Session):
"""
Test external retrieval with non-external provider (should return empty).
@ -474,15 +468,15 @@ class TestHitTestingServiceExternalRetrieve:
# Act
result = HitTestingService.external_retrieve(
dataset, query, account, external_retrieval_model, metadata_filtering_conditions, session=mock_db_session
dataset, query, account, external_retrieval_model, metadata_filtering_conditions, session=sqlite_session
)
# Assert
assert result["query"]["content"] == query
assert result["records"] == []
mock_db_session.add.assert_not_called()
assert sqlite_session.scalar(select(DatasetQuery)) is None
def test_external_retrieve_with_metadata_filtering(self, mock_db_session):
def test_external_retrieve_with_metadata_filtering(self, sqlite_session: Session):
"""
Test external retrieval with metadata filtering conditions.
@ -514,7 +508,7 @@ class TestHitTestingServiceExternalRetrieve:
account,
external_retrieval_model,
metadata_filtering_conditions,
session=mock_db_session,
session=sqlite_session,
)
# Assert
@ -523,7 +517,7 @@ class TestHitTestingServiceExternalRetrieve:
call_kwargs = mock_external_retrieve.call_args[1]
assert call_kwargs["metadata_filtering_conditions"] == metadata_filtering_conditions
def test_external_retrieve_empty_documents(self, mock_db_session):
def test_external_retrieve_empty_documents(self, sqlite_session: Session):
"""
Test external retrieval with empty document list.
@ -553,7 +547,7 @@ class TestHitTestingServiceExternalRetrieve:
account,
external_retrieval_model,
metadata_filtering_conditions,
session=mock_db_session,
session=sqlite_session,
)
# Assert
@ -569,7 +563,7 @@ class TestHitTestingServiceCompactRetrieveResponse:
ensuring documents are properly formatted into retrieval records.
"""
def test_compact_retrieve_response_success(self):
def test_compact_retrieve_response_success(self, sqlite_session: Session):
"""
Test successful response formatting.
@ -587,7 +581,6 @@ class TestHitTestingServiceCompactRetrieveResponse:
HitTestingTestDataFactory.create_retrieval_record_mock(content="Doc 1", score=0.95),
HitTestingTestDataFactory.create_retrieval_record_mock(content="Doc 2", score=0.85),
]
session = MagicMock()
with patch(
"services.hit_testing_service.RetrievalService.format_retrieval_documents", autospec=True
@ -595,7 +588,7 @@ class TestHitTestingServiceCompactRetrieveResponse:
mock_format.return_value = mock_records
# Act
result = HitTestingService.compact_retrieve_response(query, documents, session=session)
result = HitTestingService.compact_retrieve_response(query, documents, session=sqlite_session)
# Assert
assert result["query"]["content"] == query
@ -603,10 +596,11 @@ class TestHitTestingServiceCompactRetrieveResponse:
assert result["records"][0]["content"] == "Doc 1"
assert result["records"][0]["score"] == 0.95
mock_format.assert_called_once()
assert mock_format.call_args.args[0] is not session
assert mock_format.call_args.args[0] is not sqlite_session
assert mock_format.call_args.args[0].get_bind() is sqlite_session.get_bind()
assert mock_format.call_args.args[1] == documents
def test_compact_retrieve_response_empty_documents(self):
def test_compact_retrieve_response_empty_documents(self, sqlite_session: Session):
"""
Test response formatting with empty document list.
@ -616,7 +610,6 @@ class TestHitTestingServiceCompactRetrieveResponse:
# Arrange
query = "test query"
documents = []
session = MagicMock()
with patch(
"services.hit_testing_service.RetrievalService.format_retrieval_documents", autospec=True
@ -624,13 +617,14 @@ class TestHitTestingServiceCompactRetrieveResponse:
mock_format.return_value = []
# Act
result = HitTestingService.compact_retrieve_response(query, documents, session=session)
result = HitTestingService.compact_retrieve_response(query, documents, session=sqlite_session)
# Assert
assert result["query"]["content"] == query
assert result["records"] == []
mock_format.assert_called_once()
assert mock_format.call_args.args[0] is not session
assert mock_format.call_args.args[0] is not sqlite_session
assert mock_format.call_args.args[0].get_bind() is sqlite_session.get_bind()
assert mock_format.call_args.args[1] == documents

View File

@ -1,9 +1,10 @@
"""Unit tests for the Agent tool inner invoke service."""
"""Unit tests for the Agent tool inner invoke service with SQLite-backed app lookup."""
from collections.abc import Generator
from unittest.mock import MagicMock, patch
import pytest
from sqlalchemy.orm import Session
from core.tools.entities.tool_entities import ToolInvokeMessage, ToolProviderType
from core.tools.errors import (
@ -12,19 +13,44 @@ from core.tools.errors import (
ToolProviderCredentialValidationError,
ToolProviderNotFoundError,
)
from models.enums import AppStatus
from models.model import App, AppMode
from services.agent_tool_inner_service import AgentToolInnerService
from services.entities.agent_tool_inner import AgentToolInvokeRequest
from services.errors.agent_tool_inner import AgentToolInnerServiceError
TENANT_ID = "11111111-1111-1111-1111-111111111111"
OTHER_TENANT_ID = "22222222-2222-2222-2222-222222222222"
USER_ID = "33333333-3333-3333-3333-333333333333"
APP_ID = "44444444-4444-4444-4444-444444444444"
def _persist_app(sqlite_session: Session, *, tenant_id: str = TENANT_ID) -> App:
app = App(
id=APP_ID,
tenant_id=tenant_id,
name="Test App",
description="",
mode=AppMode.CHAT,
status=AppStatus.NORMAL,
enable_site=False,
enable_api=False,
max_active_requests=None,
)
sqlite_session.add(app)
sqlite_session.commit()
sqlite_session.expunge_all()
return app
def _request() -> AgentToolInvokeRequest:
return AgentToolInvokeRequest.model_validate(
{
"caller": {
"tenant_id": "tenant-1",
"user_id": "user-1",
"tenant_id": TENANT_ID,
"user_id": USER_ID,
"user_from": "account",
"app_id": "app-1",
"app_id": APP_ID,
"invoke_from": "service-api",
"conversation_id": "conversation-1",
"workflow_id": "workflow-1",
@ -53,11 +79,10 @@ def _messages() -> Generator[ToolInvokeMessage, None, None]:
)
def test_invoke_uses_agent_tool_runtime_and_returns_observation() -> None:
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_invoke_uses_agent_tool_runtime_and_returns_observation(sqlite_session: Session) -> None:
fake_tool = MagicMock()
fake_app = MagicMock(id="app-1", tenant_id="tenant-1")
session = MagicMock()
session.get.return_value = fake_app
_persist_app(sqlite_session)
with (
patch(
@ -70,7 +95,7 @@ def test_invoke_uses_agent_tool_runtime_and_returns_observation() -> None:
side_effect=lambda messages, **_kwargs: messages,
),
):
response = AgentToolInnerService().invoke(_request(), session=session)
response = AgentToolInnerService().invoke(_request(), session=sqlite_session)
assert response.observation == "ok"
assert response.metadata == {
@ -82,56 +107,58 @@ def test_invoke_uses_agent_tool_runtime_and_returns_observation() -> None:
assert agent_tool.provider_type is ToolProviderType.PLUGIN
assert agent_tool.tool_parameters == {"region": "us"}
mock_invoke.assert_called_once()
assert mock_invoke.call_args.kwargs["session"] is sqlite_session
assert sqlite_session.in_transaction()
def test_invoke_raises_app_not_found_when_session_has_no_app() -> None:
session = MagicMock()
session.get.return_value = None
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_invoke_raises_app_not_found_when_session_has_no_app(sqlite_session: Session) -> None:
with pytest.raises(AgentToolInnerServiceError) as exc_info:
AgentToolInnerService().invoke(_request(), session=session)
AgentToolInnerService().invoke(_request(), session=sqlite_session)
assert exc_info.value.error_code == "app_not_found"
assert exc_info.value.status_code == 404
assert exc_info.value.description == "App not found."
assert sqlite_session.in_transaction()
def test_invoke_raises_app_tenant_mismatch_when_app_belongs_to_other_tenant() -> None:
fake_app = MagicMock(id="app-1", tenant_id="tenant-2")
session = MagicMock()
session.get.return_value = fake_app
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_invoke_raises_app_tenant_mismatch_when_app_belongs_to_other_tenant(sqlite_session: Session) -> None:
_persist_app(sqlite_session, tenant_id=OTHER_TENANT_ID)
with pytest.raises(AgentToolInnerServiceError) as exc_info:
AgentToolInnerService().invoke(_request(), session=session)
AgentToolInnerService().invoke(_request(), session=sqlite_session)
assert exc_info.value.error_code == "app_tenant_mismatch"
assert exc_info.value.status_code == 403
assert exc_info.value.description == "App does not belong to the caller tenant."
assert sqlite_session.in_transaction()
def test_invoke_maps_tool_runtime_app_not_found_value_error_to_specific_error_code() -> None:
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_invoke_maps_tool_runtime_app_not_found_value_error_to_specific_error_code(
sqlite_session: Session,
) -> None:
fake_tool = MagicMock()
fake_app = MagicMock(id="app-1", tenant_id="tenant-1")
session = MagicMock()
session.get.return_value = fake_app
_persist_app(sqlite_session)
with (
patch("services.agent_tool_inner_service.ToolManager.get_agent_tool_runtime", return_value=fake_tool),
patch("services.agent_tool_inner_service.ToolEngine.generic_invoke", side_effect=ValueError("app not found")),
):
with pytest.raises(AgentToolInnerServiceError) as exc_info:
AgentToolInnerService().invoke(_request(), session=session)
AgentToolInnerService().invoke(_request(), session=sqlite_session)
assert exc_info.value.error_code == "app_not_found"
assert exc_info.value.status_code == 404
assert exc_info.value.description == "App not found."
assert sqlite_session.in_transaction()
def test_invoke_maps_tool_invoke_error_without_private_tool_engine_helper() -> None:
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_invoke_maps_tool_invoke_error_without_private_tool_engine_helper(sqlite_session: Session) -> None:
fake_tool = MagicMock()
fake_app = MagicMock(id="app-1", tenant_id="tenant-1")
session = MagicMock()
session.get.return_value = fake_app
_persist_app(sqlite_session)
with (
patch("services.agent_tool_inner_service.ToolManager.get_agent_tool_runtime", return_value=fake_tool),
@ -141,9 +168,10 @@ def test_invoke_maps_tool_invoke_error_without_private_tool_engine_helper() -> N
),
):
with pytest.raises(AgentToolInnerServiceError) as exc_info:
AgentToolInnerService().invoke(_request(), session=session)
AgentToolInnerService().invoke(_request(), session=sqlite_session)
assert exc_info.value.error_code == "agent_tool_invoke_failed"
assert sqlite_session.in_transaction()
@pytest.mark.parametrize(
@ -154,13 +182,17 @@ def test_invoke_maps_tool_invoke_error_without_private_tool_engine_helper() -> N
(ToolParameterValidationError("query is required"), "tool_parameters_invalid"),
],
)
def test_invoke_maps_runtime_lookup_errors_to_service_error_codes(error: Exception, expected_code: str) -> None:
fake_app = MagicMock(id="app-1", tenant_id="tenant-1")
session = MagicMock()
session.get.return_value = fake_app
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_invoke_maps_runtime_lookup_errors_to_service_error_codes(
error: Exception,
expected_code: str,
sqlite_session: Session,
) -> None:
_persist_app(sqlite_session)
with patch("services.agent_tool_inner_service.ToolManager.get_agent_tool_runtime", side_effect=error):
with pytest.raises(AgentToolInnerServiceError) as exc_info:
AgentToolInnerService().invoke(_request(), session=session)
AgentToolInnerService().invoke(_request(), session=sqlite_session)
assert exc_info.value.error_code == expected_code
assert sqlite_session.in_transaction()

View File

@ -1,12 +1,18 @@
import json
import logging
from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from sqlalchemy import select
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session
import services.async_workflow_service as async_workflow_service_module
from models.enums import AppTriggerType, CreatorUserRole, WorkflowRunTriggeredFrom, WorkflowTriggerStatus
from models.model import App, AppMode
from models.trigger import WorkflowTriggerLog
from services.async_workflow_service import AsyncWorkflowService
from services.errors.app import QuotaExceededError, WorkflowNotFoundError
from services.workflow.entities import AsyncTriggerResponse, TriggerData
@ -37,25 +43,66 @@ class AsyncWorkflowServiceTestDataFactory:
)
@staticmethod
def create_trigger_log_with_data(trigger_data: TriggerData, retry_count: int = 0) -> MagicMock:
"""Create a mock trigger log with serialized trigger data."""
trigger_log = MagicMock()
trigger_log.id = "trigger-log-123"
trigger_log.trigger_data = trigger_data.model_dump_json()
trigger_log.retry_count = retry_count
trigger_log.error = "previous-error"
trigger_log.status = WorkflowTriggerStatus.FAILED
trigger_log.to_dict.return_value = {"id": trigger_log.id}
def create_app(app_id: str = "app-123", tenant_id: str = "tenant-123") -> App:
"""Create an app that can be persisted for trigger lookup tests."""
return App(
id=app_id,
tenant_id=tenant_id,
name="Async workflow app",
description="",
mode=AppMode.WORKFLOW,
enable_site=True,
enable_api=True,
max_active_requests=0,
)
@staticmethod
def create_trigger_log_with_data(
trigger_data: TriggerData,
*,
trigger_log_id: str = "trigger-log-123",
retry_count: int = 0,
status: WorkflowTriggerStatus = WorkflowTriggerStatus.FAILED,
created_at: datetime | None = None,
) -> WorkflowTriggerLog:
"""Create a persistent trigger log with serialized trigger data."""
trigger_log = WorkflowTriggerLog(
tenant_id=trigger_data.tenant_id,
app_id=trigger_data.app_id,
workflow_id=trigger_data.workflow_id or "workflow-123",
workflow_run_id=None,
root_node_id=trigger_data.root_node_id,
trigger_metadata="{}",
trigger_type=trigger_data.trigger_type,
trigger_data=trigger_data.model_dump_json(),
inputs=json.dumps(dict(trigger_data.inputs)),
outputs=None,
status=status,
error="previous-error",
queue_name=QueuePriority.SANDBOX,
celery_task_id=None,
created_by_role=CreatorUserRole.END_USER,
created_by="end-user-123",
retry_count=retry_count,
elapsed_time=None,
total_tokens=None,
triggered_at=None,
finished_at=None,
)
trigger_log.id = trigger_log_id
if created_at is not None:
trigger_log.created_at = created_at
return trigger_log
@pytest.mark.usefixtures("sqlite_session")
@pytest.mark.parametrize("sqlite_session", [(App, WorkflowTriggerLog)], indirect=True)
class TestAsyncWorkflowService:
@pytest.fixture
def async_workflow_trigger_mocks(self):
"""Shared fixture for async workflow trigger tests.
Yields mocks for:
- repo: SQLAlchemyWorkflowTriggerLogRepository
- dispatcher_manager_class: QueueDispatcherManager class
- dispatcher: dispatcher instance
- quota_service: QuotaService mock
@ -64,23 +111,10 @@ class TestAsyncWorkflowService:
- team_task: execute_workflow_team
- sandbox_task: execute_workflow_sandbox
"""
mock_repo = MagicMock()
def _create_side_effect(new_log):
new_log.id = "trigger-log-123"
return new_log
mock_repo.create.side_effect = _create_side_effect
mock_dispatcher = MagicMock()
mock_quota_service = MagicMock()
with (
patch.object(
async_workflow_service_module,
"SQLAlchemyWorkflowTriggerLogRepository",
return_value=mock_repo,
),
patch.object(async_workflow_service_module, "QueueDispatcherManager") as mock_dispatcher_manager_class,
patch.object(async_workflow_service_module, "WorkflowService"),
patch.object(
@ -100,7 +134,6 @@ class TestAsyncWorkflowService:
mock_dispatcher_manager_class.return_value.get_dispatcher.return_value = mock_dispatcher
yield {
"repo": mock_repo,
"dispatcher_manager_class": mock_dispatcher_manager_class,
"dispatcher": mock_dispatcher,
"quota_service": mock_quota_service,
@ -119,15 +152,16 @@ class TestAsyncWorkflowService:
],
)
def test_should_dispatch_to_matching_celery_task_when_triggering_workflow(
self, queue_name, selected_task_attr, async_workflow_trigger_mocks
self,
queue_name,
selected_task_attr,
async_workflow_trigger_mocks,
sqlite_session: Session,
):
"""Test queue-based task routing and successful async trigger response."""
# Arrange
session = MagicMock()
session.commit = MagicMock()
app_model = MagicMock()
app_model.id = "app-123"
session.scalar.return_value = app_model
sqlite_session.add(AsyncWorkflowServiceTestDataFactory.create_app())
sqlite_session.commit()
trigger_data = AsyncWorkflowServiceTestDataFactory.create_trigger_data()
workflow = MagicMock()
workflow.id = "workflow-123"
@ -153,20 +187,25 @@ class TestAsyncWorkflowService:
user = DummyAccount("account-123")
# Act
result = AsyncWorkflowService.trigger_workflow_async(session=session, user=user, trigger_data=trigger_data)
result = AsyncWorkflowService.trigger_workflow_async(
session=sqlite_session, user=user, trigger_data=trigger_data
)
# Assert
assert isinstance(result, AsyncTriggerResponse)
assert result.workflow_trigger_log_id == "trigger-log-123"
assert result.workflow_trigger_log_id
assert result.task_id == "task-123"
assert result.status == "queued"
assert result.queue == queue_name
mocks["quota_service"].reserve.assert_called_once()
quota_charge_mock.commit.assert_called_once()
assert session.commit.call_count == 3
assert not sqlite_session.in_transaction()
created_log = mocks["repo"].create.call_args[0][0]
created_log = sqlite_session.scalar(
select(WorkflowTriggerLog).where(WorkflowTriggerLog.id == result.workflow_trigger_log_id)
)
assert created_log is not None
assert created_log.status == WorkflowTriggerStatus.QUEUED
assert created_log.queue_name == queue_name
assert created_log.created_by_role == CreatorUserRole.ACCOUNT
@ -182,18 +221,17 @@ class TestAsyncWorkflowService:
}
for task_attr, task_mock in task_mocks.items():
if task_attr == selected_task_attr:
task_mock.delay.assert_called_once_with({"workflow_trigger_log_id": "trigger-log-123"})
task_mock.delay.assert_called_once_with({"workflow_trigger_log_id": result.workflow_trigger_log_id})
else:
task_mock.delay.assert_not_called()
def test_should_set_end_user_role_when_triggered_by_end_user(self, async_workflow_trigger_mocks):
def test_should_set_end_user_role_when_triggered_by_end_user(
self, async_workflow_trigger_mocks, sqlite_session: Session
):
"""Test that non-account users are tracked as END_USER in trigger logs."""
# Arrange
session = MagicMock()
session.commit = MagicMock()
app_model = MagicMock()
app_model.id = "app-123"
session.scalar.return_value = app_model
sqlite_session.add(AsyncWorkflowServiceTestDataFactory.create_app())
sqlite_session.commit()
trigger_data = AsyncWorkflowServiceTestDataFactory.create_trigger_data()
workflow = MagicMock()
workflow.id = "workflow-123"
@ -208,43 +246,43 @@ class TestAsyncWorkflowService:
user = SimpleNamespace(id="end-user-123")
# Act
AsyncWorkflowService.trigger_workflow_async(session=session, user=user, trigger_data=trigger_data)
response = AsyncWorkflowService.trigger_workflow_async(
session=sqlite_session, user=user, trigger_data=trigger_data
)
# Assert
created_log = mocks["repo"].create.call_args[0][0]
created_log = sqlite_session.get(WorkflowTriggerLog, response.workflow_trigger_log_id)
assert created_log is not None
assert created_log.created_by_role == CreatorUserRole.END_USER
assert created_log.created_by == "end-user-123"
def test_should_raise_workflow_not_found_when_app_does_not_exist(self):
def test_should_raise_workflow_not_found_when_app_does_not_exist(self, sqlite_session: Session):
"""Test trigger failure when app lookup returns no result."""
# Arrange
session = MagicMock()
session.scalar.return_value = None
trigger_data = AsyncWorkflowServiceTestDataFactory.create_trigger_data(app_id="missing-app")
with (
patch.object(async_workflow_service_module, "SQLAlchemyWorkflowTriggerLogRepository"),
patch.object(async_workflow_service_module, "QueueDispatcherManager"),
patch.object(async_workflow_service_module, "WorkflowService"),
):
# Act / Assert
with pytest.raises(WorkflowNotFoundError, match="App not found: missing-app"):
AsyncWorkflowService.trigger_workflow_async(
session=session,
session=sqlite_session,
user=SimpleNamespace(id="user-123"),
trigger_data=trigger_data,
)
def test_should_mark_log_rate_limited_and_reraise_when_quota_exceeded(
self, async_workflow_trigger_mocks, caplog: pytest.LogCaptureFixture
self,
async_workflow_trigger_mocks,
caplog: pytest.LogCaptureFixture,
sqlite_session: Session,
):
"""Test quota-exceeded path updates trigger log and preserves the quota exception."""
# Arrange
session = MagicMock()
session.commit = MagicMock()
app_model = MagicMock()
app_model.id = "app-123"
session.scalar.return_value = app_model
sqlite_session.add(AsyncWorkflowServiceTestDataFactory.create_app())
sqlite_session.commit()
trigger_data = AsyncWorkflowServiceTestDataFactory.create_trigger_data()
workflow = MagicMock()
workflow.id = "workflow-123"
@ -262,7 +300,7 @@ class TestAsyncWorkflowService:
# Act / Assert
with pytest.raises(QuotaExceededError) as exc_info:
AsyncWorkflowService.trigger_workflow_async(
session=session,
session=sqlite_session,
user=SimpleNamespace(id="user-123"),
trigger_data=trigger_data,
)
@ -270,42 +308,37 @@ class TestAsyncWorkflowService:
assert exc_info.value.feature == "workflow"
assert exc_info.value.tenant_id == "tenant-123"
assert exc_info.value.required == 1
assert session.commit.call_count == 3
updated_log = mocks["repo"].update.call_args[0][0]
assert not sqlite_session.in_transaction()
updated_log = sqlite_session.scalar(select(WorkflowTriggerLog))
assert updated_log is not None
assert updated_log.status == WorkflowTriggerStatus.RATE_LIMITED
assert "Quota limit reached" in updated_log.error
assert (
"Workflow quota exceeded for tenant tenant-123, app app-123, workflow workflow-123, "
"trigger log trigger-log-123"
f"trigger log {updated_log.id}"
) in caplog.messages
mocks["professional_task"].delay.assert_not_called()
mocks["team_task"].delay.assert_not_called()
mocks["sandbox_task"].delay.assert_not_called()
def test_should_raise_when_reinvoke_target_log_does_not_exist(self):
def test_should_raise_when_reinvoke_target_log_does_not_exist(self, sqlite_session: Session):
"""Test reinvoke_trigger error path when original trigger log is missing."""
# Arrange
session = MagicMock()
repo = MagicMock()
repo.get_by_id.return_value = None
# Act / Assert
with pytest.raises(ValueError, match="Trigger log not found: missing-log"):
AsyncWorkflowService.reinvoke_trigger(
session=sqlite_session,
user=SimpleNamespace(id="user-123"),
workflow_trigger_log_id="missing-log",
)
with patch.object(async_workflow_service_module, "SQLAlchemyWorkflowTriggerLogRepository", return_value=repo):
# Act / Assert
with pytest.raises(ValueError, match="Trigger log not found: missing-log"):
AsyncWorkflowService.reinvoke_trigger(
session=session,
user=SimpleNamespace(id="user-123"),
workflow_trigger_log_id="missing-log",
)
def test_should_update_original_log_and_requeue_when_reinvoking(self):
def test_should_update_original_log_and_requeue_when_reinvoking(self, sqlite_session: Session):
"""Test reinvoke flow updates original log state and triggers a new async run."""
# Arrange
session = MagicMock()
trigger_data = AsyncWorkflowServiceTestDataFactory.create_trigger_data()
trigger_log = AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(trigger_data, retry_count=1)
repo = MagicMock()
repo.get_by_id.return_value = trigger_log
sqlite_session.add(trigger_log)
sqlite_session.commit()
expected_response = AsyncTriggerResponse(
workflow_trigger_log_id="new-trigger-log-456",
@ -315,7 +348,6 @@ class TestAsyncWorkflowService:
)
with (
patch.object(async_workflow_service_module, "SQLAlchemyWorkflowTriggerLogRepository", return_value=repo),
patch.object(
async_workflow_service_module.AsyncWorkflowService,
"trigger_workflow_async",
@ -326,145 +358,142 @@ class TestAsyncWorkflowService:
# Act
response = AsyncWorkflowService.reinvoke_trigger(
session=session,
session=sqlite_session,
user=user,
workflow_trigger_log_id="trigger-log-123",
)
# Assert
assert response == expected_response
assert not sqlite_session.in_transaction()
sqlite_session.refresh(trigger_log)
assert trigger_log.status == WorkflowTriggerStatus.RETRYING
assert trigger_log.retry_count == 2
assert trigger_log.error is None
assert trigger_log.triggered_at is not None
repo.update.assert_called_once_with(trigger_log)
session.commit.assert_called_once()
called_trigger_data = mock_trigger_workflow_async.call_args.args[1]
assert isinstance(called_trigger_data, TriggerData)
assert called_trigger_data.app_id == "app-123"
@pytest.mark.parametrize(
("repo_result", "expected"),
("lookup_id", "tenant_id", "expected_id"),
[
(None, None),
(MagicMock(), {"id": "trigger-log-123"}),
("missing-log", "tenant-123", None),
("trigger-log-123", "tenant-123", "trigger-log-123"),
("trigger-log-123", "other-tenant", None),
],
)
def test_should_return_trigger_log_dict_or_none(self, repo_result, expected):
"""Test get_trigger_log returns serialized log data or None."""
def test_should_return_trigger_log_dict_or_none(
self,
lookup_id: str,
tenant_id: str,
expected_id: str | None,
sqlite_session: Session,
sqlite_engine: Engine,
):
"""Test get_trigger_log returns persisted data with tenant isolation."""
# Arrange
mock_session = MagicMock()
mock_repo = MagicMock()
fake_engine = MagicMock()
mock_repo.get_by_id.return_value = repo_result
if repo_result:
repo_result.to_dict.return_value = expected
trigger_data = AsyncWorkflowServiceTestDataFactory.create_trigger_data()
sqlite_session.add(AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(trigger_data))
sqlite_session.commit()
mock_session_context = MagicMock()
mock_session_context.__enter__.return_value = mock_session
mock_session_context.__exit__.return_value = None
mock_sessionmaker = MagicMock()
mock_sessionmaker.return_value.begin.return_value = mock_session_context
with (
patch.object(async_workflow_service_module, "db", new=SimpleNamespace(engine=fake_engine)),
patch.object(async_workflow_service_module, "sessionmaker", mock_sessionmaker),
patch.object(
async_workflow_service_module,
"SQLAlchemyWorkflowTriggerLogRepository",
return_value=mock_repo,
),
):
with patch.object(async_workflow_service_module, "db", SimpleNamespace(engine=sqlite_engine)):
# Act
result = AsyncWorkflowService.get_trigger_log("trigger-log-123", tenant_id="tenant-123")
result = AsyncWorkflowService.get_trigger_log(lookup_id, tenant_id=tenant_id)
# Assert
assert result == expected
mock_sessionmaker.assert_called_once_with(fake_engine)
mock_repo.get_by_id.assert_called_once_with("trigger-log-123", "tenant-123")
assert (result["id"] if result else None) == expected_id
def test_should_return_recent_logs_as_dict_list(self):
"""Test get_recent_logs converts repository models into dictionaries."""
def test_should_return_recent_logs_as_dict_list(self, sqlite_session: Session, sqlite_engine: Engine):
"""Test recent logs are ordered, paginated, and tenant/app scoped."""
# Arrange
mock_session = MagicMock()
mock_repo = MagicMock()
log1 = MagicMock()
log1.to_dict.return_value = {"id": "log-1"}
log2 = MagicMock()
log2.to_dict.return_value = {"id": "log-2"}
mock_repo.get_recent_logs.return_value = [log1, log2]
now = datetime.now(UTC)
logs = [
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
AsyncWorkflowServiceTestDataFactory.create_trigger_data(),
trigger_log_id=f"log-{index}",
created_at=now - timedelta(minutes=index),
)
for index in range(1, 4)
]
logs.extend(
[
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
AsyncWorkflowServiceTestDataFactory.create_trigger_data(tenant_id="other-tenant"),
trigger_log_id="other-tenant-log",
created_at=now,
),
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
AsyncWorkflowServiceTestDataFactory.create_trigger_data(app_id="other-app"),
trigger_log_id="other-app-log",
created_at=now,
),
]
)
sqlite_session.add_all(logs)
sqlite_session.commit()
mock_session_context = MagicMock()
mock_session_context.__enter__.return_value = mock_session
mock_session_context.__exit__.return_value = None
mock_sessionmaker = MagicMock()
mock_sessionmaker.return_value.begin.return_value = mock_session_context
with (
patch.object(async_workflow_service_module, "db", new=SimpleNamespace(engine=MagicMock())),
patch.object(async_workflow_service_module, "sessionmaker", mock_sessionmaker),
patch.object(
async_workflow_service_module,
"SQLAlchemyWorkflowTriggerLogRepository",
return_value=mock_repo,
),
):
with patch.object(async_workflow_service_module, "db", SimpleNamespace(engine=sqlite_engine)):
# Act
result = AsyncWorkflowService.get_recent_logs(
tenant_id="tenant-123",
app_id="app-123",
hours=12,
limit=50,
offset=10,
limit=2,
offset=1,
)
# Assert
assert result == [{"id": "log-1"}, {"id": "log-2"}]
mock_repo.get_recent_logs.assert_called_once_with(
tenant_id="tenant-123",
app_id="app-123",
hours=12,
limit=50,
offset=10,
)
assert [log["id"] for log in result] == ["log-2", "log-3"]
def test_should_return_failed_logs_for_retry_as_dict_list(self):
"""Test get_failed_logs_for_retry serializes repository logs into dicts."""
def test_should_return_failed_logs_for_retry_as_dict_list(self, sqlite_session: Session, sqlite_engine: Engine):
"""Test retry candidates are status, retry-count, and tenant scoped."""
# Arrange
mock_session = MagicMock()
mock_repo = MagicMock()
log = MagicMock()
log.to_dict.return_value = {"id": "failed-log-1"}
mock_repo.get_failed_for_retry.return_value = [log]
mock_session_context = MagicMock()
mock_session_context.__enter__.return_value = mock_session
mock_session_context.__exit__.return_value = None
mock_sessionmaker = MagicMock()
mock_sessionmaker.return_value.begin.return_value = mock_session_context
with (
patch.object(async_workflow_service_module, "db", new=SimpleNamespace(engine=MagicMock())),
patch.object(async_workflow_service_module, "sessionmaker", mock_sessionmaker),
patch.object(
async_workflow_service_module,
"SQLAlchemyWorkflowTriggerLogRepository",
return_value=mock_repo,
now = datetime.now(UTC)
candidates = [
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
AsyncWorkflowServiceTestDataFactory.create_trigger_data(),
trigger_log_id="failed-log-1",
retry_count=1,
created_at=now - timedelta(minutes=2),
),
):
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
AsyncWorkflowServiceTestDataFactory.create_trigger_data(),
trigger_log_id="rate-limited-log",
retry_count=2,
status=WorkflowTriggerStatus.RATE_LIMITED,
created_at=now - timedelta(minutes=1),
),
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
AsyncWorkflowServiceTestDataFactory.create_trigger_data(),
trigger_log_id="retry-limit-log",
retry_count=4,
),
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
AsyncWorkflowServiceTestDataFactory.create_trigger_data(tenant_id="other-tenant"),
trigger_log_id="other-tenant-log",
),
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
AsyncWorkflowServiceTestDataFactory.create_trigger_data(),
trigger_log_id="queued-log",
status=WorkflowTriggerStatus.QUEUED,
),
]
sqlite_session.add_all(candidates)
sqlite_session.commit()
with patch.object(async_workflow_service_module, "db", SimpleNamespace(engine=sqlite_engine)):
# Act
result = AsyncWorkflowService.get_failed_logs_for_retry(tenant_id="tenant-123", max_retry_count=4, limit=20)
# Assert
assert result == [{"id": "failed-log-1"}]
mock_repo.get_failed_for_retry.assert_called_once_with(tenant_id="tenant-123", max_retry_count=4, limit=20)
assert [log["id"] for log in result] == ["failed-log-1", "rate-limited-log"]
@pytest.mark.usefixtures("sqlite_session")
@pytest.mark.parametrize("sqlite_session", [(App, WorkflowTriggerLog)], indirect=True)
class TestAsyncWorkflowServiceGetWorkflow:
def test_should_return_specific_workflow_when_workflow_id_exists(self):
def test_should_return_specific_workflow_when_workflow_id_exists(self, sqlite_session: Session):
"""Test _get_workflow returns published workflow by id when provided."""
# Arrange
workflow_service = MagicMock()
@ -473,19 +502,18 @@ class TestAsyncWorkflowServiceGetWorkflow:
workflow_service.get_published_workflow_by_id.return_value = workflow
# Act
session = MagicMock()
result = AsyncWorkflowService._get_workflow(
workflow_service, app_model, workflow_id="workflow-123", session=session
workflow_service, app_model, workflow_id="workflow-123", session=sqlite_session
)
# Assert
assert result == workflow
workflow_service.get_published_workflow_by_id.assert_called_once_with(
app_model, "workflow-123", session=session
app_model, "workflow-123", session=sqlite_session
)
workflow_service.get_published_workflow.assert_not_called()
def test_should_raise_when_specific_workflow_id_not_found(self):
def test_should_raise_when_specific_workflow_id_not_found(self, sqlite_session: Session):
"""Test _get_workflow raises WorkflowNotFoundError for unknown workflow id."""
# Arrange
workflow_service = MagicMock()
@ -495,10 +523,10 @@ class TestAsyncWorkflowServiceGetWorkflow:
# Act / Assert
with pytest.raises(WorkflowNotFoundError, match="Published workflow not found: workflow-404"):
AsyncWorkflowService._get_workflow(
workflow_service, app_model, workflow_id="workflow-404", session=MagicMock()
workflow_service, app_model, workflow_id="workflow-404", session=sqlite_session
)
def test_should_return_default_published_workflow_when_workflow_id_not_provided(self):
def test_should_return_default_published_workflow_when_workflow_id_not_provided(self, sqlite_session: Session):
"""Test _get_workflow returns default published workflow when no id is provided."""
# Arrange
workflow_service = MagicMock()
@ -508,15 +536,14 @@ class TestAsyncWorkflowServiceGetWorkflow:
workflow_service.get_published_workflow.return_value = workflow
# Act
session = MagicMock()
result = AsyncWorkflowService._get_workflow(workflow_service, app_model, session=session)
result = AsyncWorkflowService._get_workflow(workflow_service, app_model, session=sqlite_session)
# Assert
assert result == workflow
workflow_service.get_published_workflow.assert_called_once_with(app_model, session=session)
workflow_service.get_published_workflow.assert_called_once_with(app_model, session=sqlite_session)
workflow_service.get_published_workflow_by_id.assert_not_called()
def test_should_raise_when_default_published_workflow_not_found(self):
def test_should_raise_when_default_published_workflow_not_found(self, sqlite_session: Session):
"""Test _get_workflow raises WorkflowNotFoundError when app has no published workflow."""
# Arrange
workflow_service = MagicMock()
@ -526,4 +553,4 @@ class TestAsyncWorkflowServiceGetWorkflow:
# Act / Assert
with pytest.raises(WorkflowNotFoundError, match="No published workflow found for app: app-123"):
AsyncWorkflowService._get_workflow(workflow_service, app_model, session=MagicMock())
AsyncWorkflowService._get_workflow(workflow_service, app_model, session=sqlite_session)

View File

@ -613,6 +613,31 @@ class TestExternalDatasetServiceCheckEndpoint:
# Act & Assert - should not raise
ExternalDatasetService.check_endpoint_and_api_key(settings)
@patch("services.external_knowledge_service.ssrf_proxy")
def test_check_endpoint_sends_json_body(self, mock_proxy, factory: ExternalDatasetServiceTestDataFactory):
"""Regression for #39402: the validation probe must POST a JSON body matching the
External Knowledge API retrieval contract, not a body-less request that providers
such as RAGFlow reject (empty POST -> 502 ERR_ZERO_SIZE_OBJECT)."""
# Arrange
settings = {"endpoint": "https://api.example.com", "api_key": "test-key"}
mock_response = MagicMock()
mock_response.status_code = 200
mock_proxy.post.return_value = mock_response
# Act
ExternalDatasetService.check_endpoint_and_api_key(settings)
# Assert - a non-empty JSON body is sent with the JSON content type
mock_proxy.post.assert_called_once()
_, call_kwargs = mock_proxy.post.call_args
assert call_kwargs["headers"]["Content-Type"] == "application/json"
assert call_kwargs["headers"]["Authorization"] == "Bearer test-key"
sent_body = json.loads(call_kwargs["data"])
assert "knowledge_id" in sent_body
assert "query" in sent_body
assert sent_body["retrieval_setting"] == {"top_k": 1, "score_threshold": 0.0}
def test_check_endpoint_missing_endpoint_key(self, factory: ExternalDatasetServiceTestDataFactory):
"""Test validation fails when endpoint key is missing."""
# Arrange

View File

@ -0,0 +1,19 @@
import pytest
from services.feature_service import FeatureService, SystemFeatureModel
def test_system_feature_model_disables_knowledge_fs_by_default() -> None:
assert SystemFeatureModel().knowledge_fs_enabled is False
@pytest.mark.parametrize("enabled", [False, True])
def test_get_system_features_reads_knowledge_fs_flag(
monkeypatch: pytest.MonkeyPatch,
enabled: bool,
) -> None:
monkeypatch.setattr("services.feature_service.dify_config.KNOWLEDGE_FS_ENABLED", enabled)
result = FeatureService.get_system_features()
assert result.knowledge_fs_enabled is enabled

View File

@ -203,6 +203,33 @@ class TestFileService:
with pytest.raises(NotFound, match="File not found"):
file_service.get_file_base64("non_existent")
def test_get_file_presigned_url_success(self, file_service: FileService, mock_db_session):
upload_file = MagicMock(spec=UploadFile)
upload_file.key = "upload_files/tenant_id/icon.png"
upload_file.mime_type = "image/png"
mock_db_session.scalar.return_value = upload_file
with (
patch.object(dify_config, "FILES_ACCESS_TIMEOUT", 300),
patch("services.file_service.storage") as mock_storage,
):
mock_storage.generate_presigned_url.return_value = "https://s3.example.com/icon.png?signature=test"
result = file_service.get_file_presigned_url(file_id="file_id", tenant_id="tenant_id")
assert result == "https://s3.example.com/icon.png?signature=test"
mock_storage.generate_presigned_url.assert_called_once_with(
"upload_files/tenant_id/icon.png",
expires_in=300,
content_type="image/png",
)
def test_get_file_presigned_url_not_found(self, file_service: FileService, mock_db_session):
mock_db_session.scalar.return_value = None
with pytest.raises(NotFound, match="File not found"):
file_service.get_file_presigned_url(file_id="file_id", tenant_id="tenant_id")
def test_upload_text_success(self, file_service: FileService, mock_db_session):
# Setup
text = "sample text"

View File

@ -10,16 +10,21 @@ from pydantic import SecretStr
from core.helper import ssrf_proxy
from core.rbac import RBACPermission
from core.tools.errors import ToolSSRFError
from services.knowledge_fs_proxy import (
from services.knowledge_fs_operations import (
KNOWLEDGE_FS_CONSOLE_OPERATIONS,
KnowledgeFSAccessDeniedError,
KnowledgeFSConfigurationError,
KnowledgeFSMethod,
KnowledgeFSOperation,
)
from services.knowledge_fs_proxy import (
KnowledgeFSAccessDeniedError,
KnowledgeFSAuthorization,
KnowledgeFSConfigurationError,
KnowledgeFSRouteNotAllowedError,
KnowledgeFSTimeoutError,
KnowledgeFSTransportError,
authorize_knowledge_fs_request,
get_knowledge_fs_operation,
proxy_authorized_knowledge_fs_request,
proxy_knowledge_fs_request,
)
from services.knowledge_fs_proxy import (
@ -28,16 +33,178 @@ from services.knowledge_fs_proxy import (
_JWT_SECRET = "production-secret-with-at-least-32-bytes"
_HAPPY_PATH_OPERATION_IDS = (
"listKnowledgeSpaces",
"createKnowledgeSpace",
"getKnowledgeSpacesById",
"getKnowledgeSpacesByIdAccessPolicy",
"patchKnowledgeSpacesByIdAccessPolicy",
"getSourceProviders",
"getKnowledgeSpacesByIdSourceConnections",
"postKnowledgeSpacesByIdSourceConnections",
"postKnowledgeSpacesByIdSourceConnectionsByConnectionIdRefresh",
"getKnowledgeSpacesByIdSources",
"postKnowledgeSpacesByIdSources",
"postKnowledgeSpacesByIdSourcesBySourceIdCrawlPreview",
"getKnowledgeSpacesByIdSourceWorkflowsByRunId",
"getKnowledgeSpacesByIdSourceWorkflowsByRunIdPages",
"postKnowledgeSpacesByIdSourceWorkflowsByRunIdCancel",
"postKnowledgeSpacesByIdSourceWorkflowsByRunIdRetry",
"postKnowledgeSpacesByIdSourceWorkflowsByRunIdSelection",
"getKnowledgeSpacesByIdSourcesBySourceIdSyncPolicy",
"putKnowledgeSpacesByIdSourcesBySourceIdSyncPolicy",
"getKnowledgeSpacesByIdLogicalDocuments",
"getKnowledgeSpacesByIdLogicalDocumentsByDocumentId",
"getKnowledgeSpacesByIdDocumentsByDocumentIdRevisions",
"getKnowledgeSpacesByIdDocumentsByDocumentIdRevisionsByRevisionChunks",
"getKnowledgeSpacesByIdProcessingTasks",
"getKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasksByTaskIdEvents",
"deleteKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasksByTaskId",
"postKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasksByTaskIdRetry",
)
_REQUIRED_EXPANDED_HAPPY_PATH_OPERATION_IDS = {
"patchKnowledgeSpacesById",
"deleteKnowledgeSpacesById",
"getKnowledgeSpacesByIdStats",
"postKnowledgeSpacesByIdSourceConnectionsOauth",
"postSourceOauthCallback",
"getKnowledgeSpacesByIdSourceConnectionsByConnectionId",
"deleteKnowledgeSpacesByIdSourceConnectionsByConnectionId",
"getKnowledgeSpacesByIdSourcesBySourceId",
"patchKnowledgeSpacesByIdSourcesBySourceId",
"deleteKnowledgeSpacesByIdSourcesBySourceId",
"putKnowledgeSpacesByIdSourcesBySourceIdCredentials",
"deleteKnowledgeSpacesByIdSourcesBySourceIdCredentials",
"postKnowledgeSpacesByIdSourcesBySourceIdSync",
"postKnowledgeSpacesByIdSourcesBySourceIdWorkflowImports",
"getKnowledgeSpacesByIdSourcesBySourceIdPages",
"getKnowledgeSpacesByIdSourcesBySourceIdFiles",
"postKnowledgeSpacesByIdSourcesBySourceIdCrawl",
"postKnowledgeSpacesByIdSourcesBySourceIdImport",
"postKnowledgeSpacesByIdSourcesBySourceIdTest",
"postKnowledgeSpacesByIdSourcesBySourceIdImportFiles",
"postKnowledgeSpacesByIdSourcesBulk",
"getKnowledgeSpacesByIdSourceWorkflows",
"getKnowledgeSpacesByIdSourceWorkflowsByRunIdBulkItems",
"getKnowledgeSpacesByIdDocuments",
"postKnowledgeSpacesByIdDocuments",
"deleteKnowledgeSpacesByIdDocumentsBulk",
"postKnowledgeSpacesByIdDocumentsBulk",
"postKnowledgeSpacesByIdDocumentsBulkReindex",
"getKnowledgeSpacesByIdDocumentsByDocumentId",
"deleteKnowledgeSpacesByIdDocumentsByDocumentId",
"deleteKnowledgeSpacesByIdLogicalDocumentsByDocumentId",
"getKnowledgeSpacesByIdDocumentsByDocumentIdOutline",
"postKnowledgeSpacesByIdDocumentsByDocumentIdRevisionsByRevisionRollback",
"patchKnowledgeSpacesByIdDocumentsByDocumentIdMetadata",
"getKnowledgeSpacesByIdDocumentsByDocumentIdRevisionsByRevisionChunksByChunkId",
"postKnowledgeSpacesByIdDocumentsByDocumentIdRevisionsByRevisionChunksByChunkIdState",
"getKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasks",
"getKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasksByTaskId",
"getKnowledgeSpacesByIdDocumentsByDocumentIdSettings",
"putKnowledgeSpacesByIdDocumentsByDocumentIdSettings",
"getJobsById",
"deleteJobsById",
"postJobsByIdRetry",
"getDeletionJobsByJobId",
"postDeletionJobsByJobIdRetry",
"getBulkJobsById",
}
_EXPANDED_EXTERNAL_SOURCE_OPERATION_IDS = {
operation_id
for operation_id in _REQUIRED_EXPANDED_HAPPY_PATH_OPERATION_IDS
if "Source" in operation_id or operation_id == "postSourceOauthCallback"
}
_OPERATION_AUTHORIZATION_POLICIES = {
"listKnowledgeSpaces": (RBACPermission.DATASET_READONLY, "reader"),
"createKnowledgeSpace": (RBACPermission.DATASET_CREATE_AND_MANAGEMENT, "dataset_editor"),
"getKnowledgeSpacesById": (RBACPermission.DATASET_READONLY, "reader"),
"getKnowledgeSpacesByIdAccessPolicy": (RBACPermission.DATASET_READONLY, "reader"),
"patchKnowledgeSpacesByIdAccessPolicy": (RBACPermission.DATASET_ACCESS_CONFIG, "admin"),
"getSourceProviders": (RBACPermission.DATASET_EXTERNAL_CONNECT, "dataset_editor"),
"getKnowledgeSpacesByIdSourceConnections": (RBACPermission.DATASET_EXTERNAL_CONNECT, "dataset_editor"),
"postKnowledgeSpacesByIdSourceConnections": (RBACPermission.DATASET_EXTERNAL_CONNECT, "dataset_editor"),
"postKnowledgeSpacesByIdSourceConnectionsByConnectionIdRefresh": (
RBACPermission.DATASET_EXTERNAL_CONNECT,
"dataset_editor",
),
"getKnowledgeSpacesByIdSources": (RBACPermission.DATASET_READONLY, "reader"),
"postKnowledgeSpacesByIdSources": (RBACPermission.DATASET_EXTERNAL_CONNECT, "dataset_editor"),
"postKnowledgeSpacesByIdSourcesBySourceIdCrawlPreview": (
RBACPermission.DATASET_EXTERNAL_CONNECT,
"dataset_editor",
),
"getKnowledgeSpacesByIdSourceWorkflowsByRunId": (
RBACPermission.DATASET_EXTERNAL_CONNECT,
"dataset_editor",
),
"getKnowledgeSpacesByIdSourceWorkflowsByRunIdPages": (
RBACPermission.DATASET_EXTERNAL_CONNECT,
"dataset_editor",
),
"postKnowledgeSpacesByIdSourceWorkflowsByRunIdCancel": (
RBACPermission.DATASET_EXTERNAL_CONNECT,
"dataset_editor",
),
"postKnowledgeSpacesByIdSourceWorkflowsByRunIdRetry": (
RBACPermission.DATASET_EXTERNAL_CONNECT,
"dataset_editor",
),
"postKnowledgeSpacesByIdSourceWorkflowsByRunIdSelection": (
RBACPermission.DATASET_EXTERNAL_CONNECT,
"dataset_editor",
),
"getKnowledgeSpacesByIdSourcesBySourceIdSyncPolicy": (RBACPermission.DATASET_READONLY, "reader"),
"putKnowledgeSpacesByIdSourcesBySourceIdSyncPolicy": (RBACPermission.DATASET_EDIT, "dataset_editor"),
"getKnowledgeSpacesByIdLogicalDocuments": (RBACPermission.DATASET_READONLY, "reader"),
"getKnowledgeSpacesByIdLogicalDocumentsByDocumentId": (RBACPermission.DATASET_READONLY, "reader"),
"getKnowledgeSpacesByIdDocumentsByDocumentIdRevisions": (RBACPermission.DATASET_READONLY, "reader"),
"getKnowledgeSpacesByIdDocumentsByDocumentIdRevisionsByRevisionChunks": (
RBACPermission.DATASET_READONLY,
"reader",
),
"getKnowledgeSpacesByIdProcessingTasks": (RBACPermission.DATASET_READONLY, "reader"),
"getKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasksByTaskIdEvents": (
RBACPermission.DATASET_READONLY,
"reader",
),
"deleteKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasksByTaskId": (
RBACPermission.DATASET_EDIT,
"dataset_editor",
),
"postKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasksByTaskIdRetry": (
RBACPermission.DATASET_EDIT,
"dataset_editor",
),
}
def _materialized_path(operation: KnowledgeFSOperation) -> str:
segments = []
for segment in operation.path.split("/"):
if segment == "{revision}":
segments.append("1")
elif segment.startswith("{"):
segments.append("00000000-0000-4000-8000-000000000001")
else:
segments.append(segment)
return "/".join(segments)
def _set_config(
monkeypatch: pytest.MonkeyPatch,
*,
base_url: str | None = "http://knowledge-fs.test",
sse_read_timeout_seconds: float = 90.0,
timeout_seconds: float = 7.5,
jwt_secret: str | None = _JWT_SECRET,
) -> None:
values = {
"KNOWLEDGE_FS_BASE_URL": base_url,
"KNOWLEDGE_FS_SSE_READ_TIMEOUT_SECONDS": sse_read_timeout_seconds,
"KNOWLEDGE_FS_TIMEOUT_SECONDS": timeout_seconds,
"KNOWLEDGE_FS_JWT_SECRET": SecretStr(jwt_secret) if jwt_secret is not None else None,
}
@ -45,43 +212,68 @@ def _set_config(
monkeypatch.setattr(f"services.knowledge_fs_proxy.dify_config.{name}", value, raising=False)
def test_console_registry_starts_with_list_and_create_operations() -> None:
assert tuple(operation.operation_id for operation in KNOWLEDGE_FS_CONSOLE_OPERATIONS) == (
"listKnowledgeSpaces",
"createKnowledgeSpace",
def _processing_task_events_path() -> str:
operation = next(
operation
for operation in KNOWLEDGE_FS_CONSOLE_OPERATIONS
if operation.operation_id == "getKnowledgeSpacesByIdDocumentsByDocumentIdProcessingTasksByTaskIdEvents"
)
return _materialized_path(operation)
def test_console_registry_exposes_only_the_new_rag_happy_path_operations() -> None:
operation_ids = {operation.operation_id for operation in KNOWLEDGE_FS_CONSOLE_OPERATIONS}
assert operation_ids == set(_HAPPY_PATH_OPERATION_IDS) | _REQUIRED_EXPANDED_HAPPY_PATH_OPERATION_IDS
def test_console_registry_exposes_existing_upstream_contracts_needed_by_all_happy_path_pages() -> None:
operation_ids = {operation.operation_id for operation in KNOWLEDGE_FS_CONSOLE_OPERATIONS}
assert operation_ids >= _REQUIRED_EXPANDED_HAPPY_PATH_OPERATION_IDS
def test_console_registry_preserves_explicit_scope_and_authorization_policies() -> None:
for operation in KNOWLEDGE_FS_CONSOLE_OPERATIONS:
is_read = operation.method == "GET"
assert operation.required_scope == f"knowledge-spaces:{'read' if is_read else 'write'}"
expected_policy = _OPERATION_AUTHORIZATION_POLICIES.get(operation.operation_id)
if expected_policy is None and operation.operation_id in _EXPANDED_EXTERNAL_SOURCE_OPERATION_IDS:
expected_policy = (RBACPermission.DATASET_EXTERNAL_CONNECT, "dataset_editor")
if expected_policy is None:
expected_policy = (
(RBACPermission.DATASET_READONLY, "reader")
if is_read
else (RBACPermission.DATASET_EDIT, "dataset_editor")
)
assert (operation.rbac_permission, operation.legacy_role) == expected_policy
assert operation.response_headers == ("x-trace-id",)
def test_console_registry_preserves_special_transport_contracts() -> None:
crawl_preview = get_knowledge_fs_operation(
"POST",
"knowledge-spaces/00000000-0000-4000-8000-000000000001/sources/"
"00000000-0000-4000-8000-000000000002/crawl-preview",
)
selection = get_knowledge_fs_operation(
"POST",
"knowledge-spaces/00000000-0000-4000-8000-000000000001/source-workflows/"
"00000000-0000-4000-8000-000000000002/selection",
)
events = get_knowledge_fs_operation(
"GET",
"knowledge-spaces/00000000-0000-4000-8000-000000000001/documents/"
"00000000-0000-4000-8000-000000000002/processing-tasks/"
"00000000-0000-4000-8000-000000000003/events",
)
@pytest.mark.parametrize(
("method", "operation_id", "scope", "permission", "requires_dataset_editor"),
[
("GET", "listKnowledgeSpaces", "knowledge-spaces:read", RBACPermission.DATASET_READONLY, False),
(
"POST",
"createKnowledgeSpace",
"knowledge-spaces:write",
RBACPermission.DATASET_CREATE_AND_MANAGEMENT,
True,
),
],
)
def test_console_registry_preserves_contract_and_policy(
method: KnowledgeFSMethod,
operation_id: str,
scope: str,
permission: RBACPermission,
requires_dataset_editor: bool,
) -> None:
operation = get_knowledge_fs_operation(method, "knowledge-spaces")
assert operation.operation_id == operation_id
assert operation.required_scope == scope
assert operation.rbac_permission == permission
assert operation.requires_dataset_editor is requires_dataset_editor
assert operation.max_response_bytes == 1_048_576
assert operation.request_headers == ("x-trace-id",)
assert operation.response_headers == ("x-trace-id",)
assert operation.response_media_types == ("application/json",)
assert crawl_preview.request_headers == ("idempotency-key", "x-trace-id")
assert selection.request_headers == ("idempotency-key", "x-trace-id")
assert events.response_kind == "stream"
assert events.max_response_bytes == 67_108_864
assert events.request_headers == ("last-event-id", "x-trace-id")
assert events.response_media_types == ("text/event-stream",)
def test_unconfigured_kfs_is_rejected_before_external_io(monkeypatch: pytest.MonkeyPatch) -> None:
@ -148,7 +340,112 @@ def test_proxy_forwards_only_registry_declared_headers(monkeypatch: pytest.Monke
assert forward.call_args.kwargs["request_headers"] == {"x-trace-id": "trace-1"}
def test_authorization_rejects_workspace_rbac_denial(monkeypatch: pytest.MonkeyPatch) -> None:
def test_authorized_proxy_does_not_repeat_workspace_rbac(monkeypatch: pytest.MonkeyPatch) -> None:
account = MagicMock(id="account-1", is_dataset_editor=True)
check_access = MagicMock(return_value=True)
forward = MagicMock(return_value=MagicMock())
monkeypatch.setattr("services.knowledge_fs_proxy.RBACService.CheckAccess.check", check_access)
monkeypatch.setattr("services.knowledge_fs_proxy._forward_knowledge_fs_request", forward)
operation = get_knowledge_fs_operation("POST", "knowledge-spaces")
authorization = authorize_knowledge_fs_request(
account=account,
tenant_id="tenant-1",
method=operation.method,
path=_materialized_path(operation),
)
proxy_authorized_knowledge_fs_request(authorization=authorization)
check_access.assert_called_once()
assert forward.call_args.kwargs["account_id"] == "account-1"
assert forward.call_args.kwargs["tenant_id"] == "tenant-1"
assert forward.call_args.kwargs["method"] == "POST"
assert forward.call_args.kwargs["path"] == "knowledge-spaces"
def test_authorization_capability_cannot_be_constructed_directly() -> None:
operation = get_knowledge_fs_operation("POST", "knowledge-spaces")
with pytest.raises(KnowledgeFSAccessDeniedError, match="must be created by workspace authorization"):
KnowledgeFSAuthorization("account-1", "tenant-1", operation)
def test_authorization_resolves_the_canonical_operation_policy(monkeypatch: pytest.MonkeyPatch) -> None:
account = MagicMock(id="account-1", is_dataset_editor=False, is_admin_or_owner=False)
check_access = MagicMock(return_value=True)
monkeypatch.setattr("services.knowledge_fs_proxy.RBACService.CheckAccess.check", check_access)
with pytest.raises(KnowledgeFSAccessDeniedError, match="dataset edit access"):
authorize_knowledge_fs_request(
account=account,
tenant_id="tenant-1",
method="POST",
path="knowledge-spaces",
)
check_access.assert_not_called()
@pytest.mark.parametrize(
("attribute", "value"),
[
("account_id", "account-2"),
("tenant_id", "tenant-2"),
("operation", get_knowledge_fs_operation("GET", "knowledge-spaces")),
],
)
def test_authorization_capability_binding_cannot_be_mutated(
monkeypatch: pytest.MonkeyPatch,
attribute: str,
value: object,
) -> None:
account = MagicMock(id="account-1", is_dataset_editor=True)
monkeypatch.setattr(
"services.knowledge_fs_proxy.RBACService.CheckAccess.check",
MagicMock(return_value=True),
)
authorization = authorize_knowledge_fs_request(
account=account,
tenant_id="tenant-1",
method="POST",
path="knowledge-spaces",
)
with pytest.raises(AttributeError):
setattr(authorization, attribute, value)
assert authorization.account_id == "account-1"
assert authorization.tenant_id == "tenant-1"
assert authorization.operation == get_knowledge_fs_operation("POST", "knowledge-spaces")
def test_authorization_capability_cannot_be_reused(monkeypatch: pytest.MonkeyPatch) -> None:
account = MagicMock(id="account-1", is_dataset_editor=True)
forward = MagicMock(return_value=MagicMock())
monkeypatch.setattr(
"services.knowledge_fs_proxy.RBACService.CheckAccess.check",
MagicMock(return_value=True),
)
monkeypatch.setattr("services.knowledge_fs_proxy._forward_knowledge_fs_request", forward)
authorization = authorize_knowledge_fs_request(
account=account,
tenant_id="tenant-1",
method="POST",
path="knowledge-spaces",
)
proxy_authorized_knowledge_fs_request(authorization=authorization)
with pytest.raises(KnowledgeFSAccessDeniedError, match="already been used"):
proxy_authorized_knowledge_fs_request(authorization=authorization)
forward.assert_called_once()
@pytest.mark.parametrize("operation", KNOWLEDGE_FS_CONSOLE_OPERATIONS, ids=lambda operation: operation.operation_id)
def test_authorization_rejects_workspace_rbac_denial(
monkeypatch: pytest.MonkeyPatch,
operation: KnowledgeFSOperation,
) -> None:
account = MagicMock(id="account-1", is_dataset_editor=True)
check_access = MagicMock(return_value=False)
monkeypatch.setattr("services.knowledge_fs_proxy.RBACService.CheckAccess.check", check_access)
@ -157,19 +454,28 @@ def test_authorization_rejects_workspace_rbac_denial(monkeypatch: pytest.MonkeyP
authorize_knowledge_fs_request(
account=account,
tenant_id="tenant-1",
operation=get_knowledge_fs_operation("GET", "knowledge-spaces"),
method=operation.method,
path=_materialized_path(operation),
)
check_access.assert_called_once_with(
"tenant-1",
"account-1",
scene="dataset_readonly",
scene=operation.rbac_permission.value,
resource_type="dataset",
)
def test_create_rejects_non_dataset_editor_before_rbac(monkeypatch: pytest.MonkeyPatch) -> None:
account = MagicMock(id="account-1", is_dataset_editor=False)
@pytest.mark.parametrize(
"operation",
tuple(operation for operation in KNOWLEDGE_FS_CONSOLE_OPERATIONS if operation.legacy_role == "dataset_editor"),
ids=lambda operation: operation.operation_id,
)
def test_dataset_editor_operations_reject_legacy_viewers_before_rbac(
monkeypatch: pytest.MonkeyPatch,
operation: KnowledgeFSOperation,
) -> None:
account = MagicMock(id="account-1", is_dataset_editor=False, is_admin_or_owner=False)
check_access = MagicMock(return_value=True)
monkeypatch.setattr("services.knowledge_fs_proxy.RBACService.CheckAccess.check", check_access)
@ -177,41 +483,66 @@ def test_create_rejects_non_dataset_editor_before_rbac(monkeypatch: pytest.Monke
authorize_knowledge_fs_request(
account=account,
tenant_id="tenant-1",
operation=get_knowledge_fs_operation("POST", "knowledge-spaces"),
method=operation.method,
path=_materialized_path(operation),
)
check_access.assert_not_called()
def test_authorization_uses_the_declared_editor_policy(monkeypatch: pytest.MonkeyPatch) -> None:
account = MagicMock(id="account-1", is_dataset_editor=False)
def test_admin_operation_rejects_legacy_editors_before_rbac(monkeypatch: pytest.MonkeyPatch) -> None:
account = MagicMock(id="account-1", is_dataset_editor=True, is_admin_or_owner=False)
check_access = MagicMock(return_value=True)
monkeypatch.setattr("services.knowledge_fs_proxy.RBACService.CheckAccess.check", check_access)
operation = get_knowledge_fs_operation("POST", "knowledge-spaces")._replace(requires_dataset_editor=False)
operation = get_knowledge_fs_operation(
"PATCH", "knowledge-spaces/00000000-0000-4000-8000-000000000001/access-policy"
)
authorize_knowledge_fs_request(account=account, tenant_id="tenant-1", operation=operation)
with pytest.raises(KnowledgeFSAccessDeniedError, match="administration access"):
authorize_knowledge_fs_request(
account=account,
tenant_id="tenant-1",
method=operation.method,
path=_materialized_path(operation),
)
check_access.assert_not_called()
def test_authorization_uses_the_declared_reader_policy(monkeypatch: pytest.MonkeyPatch) -> None:
account = MagicMock(id="account-1", is_dataset_editor=False, is_admin_or_owner=False)
check_access = MagicMock(return_value=True)
monkeypatch.setattr("services.knowledge_fs_proxy.RBACService.CheckAccess.check", check_access)
operation = get_knowledge_fs_operation("GET", "knowledge-spaces")
authorize_knowledge_fs_request(
account=account,
tenant_id="tenant-1",
method=operation.method,
path=_materialized_path(operation),
)
check_access.assert_called_once()
@pytest.mark.parametrize(
("method", "expected_scope"),
[("GET", "knowledge-spaces:read"), ("POST", "knowledge-spaces:write")],
)
@pytest.mark.parametrize("operation", KNOWLEDGE_FS_CONSOLE_OPERATIONS, ids=lambda operation: operation.operation_id)
def test_auth_signs_current_principals_and_declared_scope(
monkeypatch: pytest.MonkeyPatch,
method: KnowledgeFSMethod,
expected_scope: str,
operation: KnowledgeFSOperation,
) -> None:
_set_config(monkeypatch)
response = httpx.Response(200, content=b'{"items":[]}', headers={"Content-Type": "application/json"})
response = httpx.Response(
200,
content=b"data" if operation.response_kind == "stream" else b'{"items":[]}',
headers={"Content-Type": operation.response_media_types[0]},
)
request = MagicMock(return_value=response)
monkeypatch.setattr("services.knowledge_fs_proxy.ssrf_proxy.make_request", request)
forward_knowledge_fs_request(
account_id="account-1",
method=method,
path="knowledge-spaces",
method=operation.method,
path=_materialized_path(operation),
tenant_id="tenant-1",
)
@ -226,7 +557,7 @@ def test_auth_signs_current_principals_and_declared_scope(
assert claims["dify_account_id"] == "dify-account:account-1"
assert claims["sub"] == "dify-workspace:tenant-1"
assert claims["tenant_id"] == "tenant-1"
assert claims["scopes"] == [expected_scope]
assert claims["scopes"] == [operation.required_scope]
assert claims["caller_kind"] == "interactive"
assert claims["exp"] - claims["iat"] == 60
@ -244,6 +575,69 @@ def test_buffered_response_rejects_non_empty_body_without_content_type(monkeypat
assert response.is_closed
def test_sse_response_remains_streaming_and_uses_the_dedicated_read_timeout(
monkeypatch: pytest.MonkeyPatch,
) -> None:
_set_config(monkeypatch, sse_read_timeout_seconds=120.0)
request = httpx.Request(
"GET",
"http://knowledge-fs.test/events",
extensions={"timeout": {"connect": 7.5, "read": 7.5}},
)
response = httpx.Response(
200,
headers={"Content-Type": "text/event-stream"},
request=request,
stream=httpx.ByteStream(b"event: progress\n\n"),
)
monkeypatch.setattr("services.knowledge_fs_proxy.ssrf_proxy.make_request", MagicMock(return_value=response))
buffer_response = MagicMock()
monkeypatch.setattr("services.knowledge_fs_proxy.ssrf_proxy.buffer_response", buffer_response)
result = forward_knowledge_fs_request(
account_id="account-dev",
method="GET",
path=_processing_task_events_path(),
tenant_id="tenant-dev",
)
assert result.response is response
assert result.response_kind == "stream"
assert request.extensions["timeout"]["read"] == 120.0
assert not response.is_closed
buffer_response.assert_not_called()
@pytest.mark.parametrize(
("headers", "message"),
[
({"Content-Type": "application/json"}, "unsupported media type"),
(
{"Content-Type": "text/event-stream", "Content-Encoding": "gzip"},
"unsupported encoding",
),
],
)
def test_sse_response_rejects_invalid_stream_headers_and_closes_upstream(
monkeypatch: pytest.MonkeyPatch,
headers: dict[str, str],
message: str,
) -> None:
_set_config(monkeypatch)
response = httpx.Response(200, headers=headers, stream=httpx.ByteStream(b"data"))
monkeypatch.setattr("services.knowledge_fs_proxy.ssrf_proxy.make_request", MagicMock(return_value=response))
with pytest.raises(KnowledgeFSTransportError, match=message):
forward_knowledge_fs_request(
account_id="account-dev",
method="GET",
path=_processing_task_events_path(),
tenant_id="tenant-dev",
)
assert response.is_closed
@pytest.mark.parametrize(
("error", "expected_exception"),
[
@ -277,10 +671,10 @@ def test_transport_failures_are_normalized(
("method", "path"),
[
("GET", "openapi.json"),
("GET", "knowledge-spaces/space-1"),
("GET", "knowledge-spaces/space-1/manifest"),
("PATCH", "knowledge-spaces"),
("POST", "queries"),
("POST", "knowledge-spaces/space-1/documents"),
("POST", "knowledge-spaces/space-1/uploads"),
],
)
def test_unregistered_route_is_rejected_before_external_io(

View File

@ -1419,6 +1419,8 @@ class TestWorkflowService:
# ===========================================================================
@pytest.mark.usefixtures("sqlite_session")
@pytest.mark.parametrize("sqlite_session", [(BuiltinToolProvider,)], indirect=True)
class TestWorkflowServiceCredentialValidation:
"""
Tests for the private credential-validation helpers on WorkflowService.
@ -1444,7 +1446,7 @@ class TestWorkflowServiceCredentialValidation:
# --- _validate_workflow_credentials: tool node (with credential_id) ---
def test_validate_workflow_credentials_should_check_tool_credential_when_credential_id_present(
self, service: WorkflowService
self, service: WorkflowService, sqlite_session: Session
) -> None:
# Arrange
nodes = [
@ -1462,11 +1464,11 @@ class TestWorkflowServiceCredentialValidation:
# Act + Assert
with patch("core.helper.credential_utils.check_credential_policy_compliance") as mock_check:
# Should not raise; mock allows the call
service._validate_workflow_credentials(workflow, session=MagicMock())
service._validate_workflow_credentials(workflow, session=sqlite_session)
mock_check.assert_called_once()
def test_validate_workflow_credentials_should_check_default_credential_when_no_credential_id(
self, service: WorkflowService
self, service: WorkflowService, sqlite_session: Session
) -> None:
# Arrange
nodes = [
@ -1483,14 +1485,13 @@ class TestWorkflowServiceCredentialValidation:
# Act
with patch.object(service, "_check_default_tool_credential") as mock_default:
session = MagicMock()
service._validate_workflow_credentials(workflow, session=session)
service._validate_workflow_credentials(workflow, session=sqlite_session)
# Assert
mock_default.assert_called_once_with("tenant-1", "my-provider", session=session)
mock_default.assert_called_once_with("tenant-1", "my-provider", session=sqlite_session)
def test_validate_workflow_credentials_should_skip_tool_node_without_provider(
self, service: WorkflowService
self, service: WorkflowService, sqlite_session: Session
) -> None:
"""Tool nodes without a provider_id should be silently skipped."""
# Arrange
@ -1499,11 +1500,11 @@ class TestWorkflowServiceCredentialValidation:
# Act + Assert (no error raised)
with patch.object(service, "_check_default_tool_credential") as mock_default:
service._validate_workflow_credentials(workflow, session=MagicMock())
service._validate_workflow_credentials(workflow, session=sqlite_session)
mock_default.assert_not_called()
def test_validate_workflow_credentials_should_validate_llm_node_with_model_config(
self, service: WorkflowService
self, service: WorkflowService, sqlite_session: Session
) -> None:
# Arrange
nodes = [
@ -1522,13 +1523,13 @@ class TestWorkflowServiceCredentialValidation:
patch.object(service, "_validate_llm_model_config") as mock_llm,
patch.object(service, "_validate_load_balancing_credentials"),
):
service._validate_workflow_credentials(workflow, session=MagicMock())
service._validate_workflow_credentials(workflow, session=sqlite_session)
# Assert
mock_llm.assert_called_once_with("tenant-1", "openai", "gpt-4")
def test_validate_workflow_credentials_should_raise_for_llm_node_missing_model(
self, service: WorkflowService
self, service: WorkflowService, sqlite_session: Session
) -> None:
"""LLM nodes without provider AND name should raise ValueError."""
# Arrange
@ -1542,10 +1543,10 @@ class TestWorkflowServiceCredentialValidation:
# Act + Assert
with pytest.raises(ValueError, match="Missing provider or model configuration"):
service._validate_workflow_credentials(workflow, session=MagicMock())
service._validate_workflow_credentials(workflow, session=sqlite_session)
def test_validate_workflow_credentials_should_wrap_unexpected_exception_in_value_error(
self, service: WorkflowService
self, service: WorkflowService, sqlite_session: Session
) -> None:
"""Non-ValueError exceptions from validation must be re-raised as ValueError."""
# Arrange
@ -1563,9 +1564,11 @@ class TestWorkflowServiceCredentialValidation:
# Act + Assert
with patch.object(service, "_validate_llm_model_config", side_effect=RuntimeError("boom")):
with pytest.raises(ValueError, match="boom"):
service._validate_workflow_credentials(workflow, session=MagicMock())
service._validate_workflow_credentials(workflow, session=sqlite_session)
def test_validate_workflow_credentials_should_validate_agent_node_model(self, service: WorkflowService) -> None:
def test_validate_workflow_credentials_should_validate_agent_node_model(
self, service: WorkflowService, sqlite_session: Session
) -> None:
# Arrange
nodes = [
{
@ -1586,12 +1589,14 @@ class TestWorkflowServiceCredentialValidation:
patch.object(service, "_validate_llm_model_config") as mock_llm,
patch.object(service, "_validate_load_balancing_credentials"),
):
service._validate_workflow_credentials(workflow, session=MagicMock())
service._validate_workflow_credentials(workflow, session=sqlite_session)
# Assert
mock_llm.assert_called_once_with("tenant-1", "openai", "gpt-4")
def test_validate_workflow_credentials_should_validate_agent_tools(self, service: WorkflowService) -> None:
def test_validate_workflow_credentials_should_validate_agent_tools(
self, service: WorkflowService, sqlite_session: Session
) -> None:
"""Each agent tool with a provider should be checked for credential compliance."""
# Arrange
nodes = [
@ -1618,12 +1623,11 @@ class TestWorkflowServiceCredentialValidation:
patch("core.helper.credential_utils.check_credential_policy_compliance") as mock_check,
patch.object(service, "_check_default_tool_credential") as mock_default,
):
session = MagicMock()
service._validate_workflow_credentials(workflow, session=session)
service._validate_workflow_credentials(workflow, session=sqlite_session)
# Assert
mock_check.assert_called_once() # provider-a has credential_id
mock_default.assert_called_once_with("tenant-1", "provider-b", session=session)
mock_default.assert_called_once_with("tenant-1", "provider-b", session=sqlite_session)
# --- _validate_llm_model_config ---
@ -1676,14 +1680,12 @@ class TestWorkflowServiceCredentialValidation:
# --- _check_default_tool_credential ---
@pytest.mark.parametrize("sqlite_session", [(BuiltinToolProvider,)], indirect=True)
def test_check_default_tool_credential_should_silently_pass_when_no_provider_found(
self, service: WorkflowService, sqlite_session: Session
) -> None:
"""Missing BuiltinToolProvider → plugin requires no credentials → no error."""
service._check_default_tool_credential("tenant-1", "some-provider", session=sqlite_session)
@pytest.mark.parametrize("sqlite_session", [(BuiltinToolProvider,)], indirect=True)
def test_check_default_tool_credential_should_raise_when_compliance_fails(
self, service: WorkflowService, sqlite_session: Session
) -> None:
@ -1746,7 +1748,9 @@ class TestWorkflowServiceCredentialValidation:
# --- _get_load_balancing_configs ---
def test_get_load_balancing_configs_should_return_empty_list_on_exception(self, service: WorkflowService) -> None:
def test_get_load_balancing_configs_should_return_empty_list_on_exception(
self, service: WorkflowService, sqlite_session: Session
) -> None:
"""Any exception during LB config retrieval should return an empty list."""
# Arrange
with patch(
@ -1754,12 +1758,14 @@ class TestWorkflowServiceCredentialValidation:
side_effect=RuntimeError("fail"),
):
# Act
result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4", session=MagicMock())
result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4", session=sqlite_session)
# Assert
assert result == []
def test_get_load_balancing_configs_should_merge_predefined_and_custom(self, service: WorkflowService) -> None:
def test_get_load_balancing_configs_should_merge_predefined_and_custom(
self, service: WorkflowService, sqlite_session: Session
) -> None:
# Arrange
predefined = [{"credential_id": "cred-a"}, {"credential_id": None}]
custom = [{"credential_id": "cred-b"}]
@ -1771,7 +1777,7 @@ class TestWorkflowServiceCredentialValidation:
],
):
# Act
result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4", session=MagicMock())
result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4", session=sqlite_session)
# Assert — only entries with a credential_id should be returned
assert len(result) == 2
@ -1780,7 +1786,7 @@ class TestWorkflowServiceCredentialValidation:
# --- _validate_load_balancing_credentials ---
def test_validate_load_balancing_credentials_should_skip_when_no_model_config(
self, service: WorkflowService
self, service: WorkflowService, sqlite_session: Session
) -> None:
"""Missing provider or model in node_data should be a no-op."""
# Arrange
@ -1788,10 +1794,10 @@ class TestWorkflowServiceCredentialValidation:
node_data: dict[str, Any] = {} # no model key
# Act + Assert (no error expected)
service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=MagicMock())
service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=sqlite_session)
def test_validate_load_balancing_credentials_should_skip_when_lb_not_enabled(
self, service: WorkflowService
self, service: WorkflowService, sqlite_session: Session
) -> None:
# Arrange
workflow = self._make_workflow([])
@ -1799,10 +1805,10 @@ class TestWorkflowServiceCredentialValidation:
# Act + Assert (no error expected)
with patch.object(service, "_is_load_balancing_enabled", return_value=False):
service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=MagicMock())
service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=sqlite_session)
def test_validate_load_balancing_credentials_should_raise_when_compliance_fails(
self, service: WorkflowService
self, service: WorkflowService, sqlite_session: Session
) -> None:
# Arrange
workflow = self._make_workflow([])
@ -1819,7 +1825,7 @@ class TestWorkflowServiceCredentialValidation:
),
):
with pytest.raises(ValueError, match="Invalid load balancing credentials"):
service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=MagicMock())
service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=sqlite_session)
# ===========================================================================

View File

@ -9,30 +9,29 @@ from services.tools.tools_transform_service import ToolTransformService
MODULE = "services.tools.tools_transform_service"
def _parameter(
name: str,
label: str,
form: ToolParameter.ToolParameterForm = ToolParameter.ToolParameterForm.FORM,
) -> ToolParameter:
return ToolParameter(
name=name,
label=I18nObject(en_US=label),
human_description=I18nObject(en_US=label),
type=ToolParameter.ToolParameterType.STRING,
form=form,
)
class TestToolTransformService:
"""Test cases for ToolTransformService.convert_tool_entity_to_api_entity method"""
def test_convert_tool_with_parameter_override(self):
"""Test that runtime parameters correctly override base parameters"""
# Create mock base parameters
base_param1 = Mock(spec=ToolParameter)
base_param1.name = "param1"
base_param1.form = ToolParameter.ToolParameterForm.FORM
base_param1.type = "string"
base_param1.label = "Base Param 1"
base_param1 = _parameter("param1", "Base Param 1")
base_param2 = _parameter("param2", "Base Param 2")
base_param2 = Mock(spec=ToolParameter)
base_param2.name = "param2"
base_param2.form = ToolParameter.ToolParameterForm.FORM
base_param2.type = "string"
base_param2.label = "Base Param 2"
# Create mock runtime parameters that override base parameters
runtime_param1 = Mock(spec=ToolParameter)
runtime_param1.name = "param1"
runtime_param1.form = ToolParameter.ToolParameterForm.FORM
runtime_param1.type = "string"
runtime_param1.label = "Runtime Param 1" # Different label to verify override
runtime_param1 = _parameter("param1", "Runtime Param 1")
# Create mock tool
mock_tool = Mock(spec=Tool)
@ -63,34 +62,19 @@ class TestToolTransformService:
# Find the overridden parameter
overridden_param = next((p for p in result.parameters if p.name == "param1"), None)
assert overridden_param is not None
assert overridden_param.label == "Runtime Param 1" # Should be runtime version
assert overridden_param.label.en_US == "Runtime Param 1" # Should be runtime version
# Find the non-overridden parameter
original_param = next((p for p in result.parameters if p.name == "param2"), None)
assert original_param is not None
assert original_param.label == "Base Param 2" # Should be base version
assert original_param.label.en_US == "Base Param 2" # Should be base version
def test_convert_tool_with_additional_runtime_parameters(self):
"""Test that additional runtime parameters are added to the final list"""
# Create mock base parameters
base_param1 = Mock(spec=ToolParameter)
base_param1.name = "param1"
base_param1.form = ToolParameter.ToolParameterForm.FORM
base_param1.type = "string"
base_param1.label = "Base Param 1"
base_param1 = _parameter("param1", "Base Param 1")
# Create mock runtime parameters - one that overrides and one that's new
runtime_param1 = Mock(spec=ToolParameter)
runtime_param1.name = "param1"
runtime_param1.form = ToolParameter.ToolParameterForm.FORM
runtime_param1.type = "string"
runtime_param1.label = "Runtime Param 1"
runtime_param2 = Mock(spec=ToolParameter)
runtime_param2.name = "runtime_only"
runtime_param2.form = ToolParameter.ToolParameterForm.FORM
runtime_param2.type = "string"
runtime_param2.label = "Runtime Only Param"
runtime_param1 = _parameter("param1", "Runtime Param 1")
runtime_param2 = _parameter("runtime_only", "Runtime Only Param")
# Create mock tool
mock_tool = Mock(spec=Tool)
@ -124,34 +108,19 @@ class TestToolTransformService:
# Verify the overridden parameter has runtime version
overridden_param = next((p for p in result.parameters if p.name == "param1"), None)
assert overridden_param is not None
assert overridden_param.label == "Runtime Param 1"
assert overridden_param.label.en_US == "Runtime Param 1"
# Verify the new runtime parameter is included
new_param = next((p for p in result.parameters if p.name == "runtime_only"), None)
assert new_param is not None
assert new_param.label == "Runtime Only Param"
assert new_param.label.en_US == "Runtime Only Param"
def test_convert_tool_with_non_form_runtime_parameters(self):
"""Test that non-FORM runtime parameters are not added as new parameters"""
# Create mock base parameters
base_param1 = Mock(spec=ToolParameter)
base_param1.name = "param1"
base_param1.form = ToolParameter.ToolParameterForm.FORM
base_param1.type = "string"
base_param1.label = "Base Param 1"
base_param1 = _parameter("param1", "Base Param 1")
# Create mock runtime parameters with different forms
runtime_param1 = Mock(spec=ToolParameter)
runtime_param1.name = "param1"
runtime_param1.form = ToolParameter.ToolParameterForm.FORM
runtime_param1.type = "string"
runtime_param1.label = "Runtime Param 1"
runtime_param2 = Mock(spec=ToolParameter)
runtime_param2.name = "llm_param"
runtime_param2.form = ToolParameter.ToolParameterForm.LLM
runtime_param2.type = "string"
runtime_param2.label = "LLM Param"
runtime_param1 = _parameter("param1", "Runtime Param 1")
runtime_param2 = _parameter("llm_param", "LLM Param", ToolParameter.ToolParameterForm.LLM)
# Create mock tool
mock_tool = Mock(spec=Tool)
@ -236,38 +205,15 @@ class TestToolTransformService:
def test_convert_tool_parameter_order_preserved(self):
"""Test that parameter order is preserved correctly"""
# Create mock base parameters in specific order
base_param1 = Mock(spec=ToolParameter)
base_param1.name = "param1"
base_param1.form = ToolParameter.ToolParameterForm.FORM
base_param1.type = "string"
base_param1.label = "Base Param 1"
base_param2 = Mock(spec=ToolParameter)
base_param2.name = "param2"
base_param2.form = ToolParameter.ToolParameterForm.FORM
base_param2.type = "string"
base_param2.label = "Base Param 2"
base_param3 = Mock(spec=ToolParameter)
base_param3.name = "param3"
base_param3.form = ToolParameter.ToolParameterForm.FORM
base_param3.type = "string"
base_param3.label = "Base Param 3"
base_param1 = _parameter("param1", "Base Param 1")
base_param2 = _parameter("param2", "Base Param 2")
base_param3 = _parameter("param3", "Base Param 3")
# Create runtime parameter that overrides middle parameter
runtime_param2 = Mock(spec=ToolParameter)
runtime_param2.name = "param2"
runtime_param2.form = ToolParameter.ToolParameterForm.FORM
runtime_param2.type = "string"
runtime_param2.label = "Runtime Param 2"
runtime_param2 = _parameter("param2", "Runtime Param 2")
# Create new runtime parameter
runtime_param4 = Mock(spec=ToolParameter)
runtime_param4.name = "param4"
runtime_param4.form = ToolParameter.ToolParameterForm.FORM
runtime_param4.type = "string"
runtime_param4.label = "Runtime Param 4"
runtime_param4 = _parameter("param4", "Runtime Param 4")
# Create mock tool
mock_tool = Mock(spec=Tool)
@ -300,7 +246,7 @@ class TestToolTransformService:
# Verify that param2 was overridden with runtime version
param2 = result.parameters[1]
assert param2.name == "param2"
assert param2.label == "Runtime Param 2"
assert param2.label.en_US == "Runtime Param 2"
class TestWorkflowProviderToUserProvider:

View File

@ -10,6 +10,8 @@ from typing import Any, cast
from unittest.mock import MagicMock
import pytest
from sqlalchemy import event
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, sessionmaker
from core.app.app_config.entities import WorkflowUIBasedAppConfig
@ -21,8 +23,9 @@ from core.app.layers.pause_state_persist_layer import (
)
from graphon.enums import WorkflowExecutionStatus
from graphon.runtime import GraphRuntimeState, VariablePool
from models.enums import CreatorUserRole
from models.model import AppMode
from models.base import TypeBase
from models.enums import CreatorUserRole, MessageStatus
from models.model import AppMode, Message
from models.workflow import WorkflowRun
from repositories.entities.workflow_pause import WorkflowPauseEntity
from services import workflow_event_snapshot_service as service_module
@ -79,23 +82,47 @@ def _build_resumption_context(task_id: str) -> WorkflowResumptionContext:
)
class _SessionContext:
def __init__(self, session: Any) -> None:
self._session = session
def __enter__(self) -> Any:
return self._session
def __exit__(self, exc_type: Any, exc: Any, tb: Any) -> bool:
return False
@pytest.fixture
def message_session_maker(sqlite_engine: Engine) -> sessionmaker[Session]:
"""Create real sessions containing only workflow messages."""
TypeBase.metadata.create_all(sqlite_engine, tables=[TypeBase.metadata.tables[Message.__tablename__]])
return sessionmaker(bind=sqlite_engine, expire_on_commit=False)
class _SessionMaker:
def __init__(self, session: Any) -> None:
self._session = session
def __call__(self) -> _SessionContext:
return _SessionContext(self._session)
def _persist_message(session_maker: sessionmaker[Session]) -> Message:
message = Message(
app_id="app-1",
model_provider="provider",
model_id="model",
override_model_configs=None,
conversation_id="conv-1",
inputs={},
query="hello",
message="",
message_tokens=0,
message_unit_price=0,
message_price_unit=0,
answer="answer",
answer_tokens=0,
answer_unit_price=0,
answer_price_unit=0,
parent_message_id=None,
provider_response_latency=0,
total_price=0,
currency="USD",
invoke_from=InvokeFrom.WEB_APP,
from_source="api",
from_end_user_id="user-1",
from_account_id=None,
app_mode=AppMode.WORKFLOW,
status=MessageStatus.NORMAL,
workflow_run_id="run-1",
)
message.id = "msg-1"
with session_maker() as session:
session.add(message)
session.commit()
return message
class _SubscriptionContext:
@ -150,12 +177,11 @@ class _PauseEntity(WorkflowPauseEntity):
class TestWorkflowEventSnapshotHelpers:
def test_get_message_context_by_conversation_should_return_none_when_no_message(self) -> None:
session = SimpleNamespace(scalar=MagicMock(return_value=None))
session_maker = _SessionMaker(session)
def test_get_message_context_by_conversation_should_return_none_when_no_message(
self, message_session_maker: sessionmaker[Session]
) -> None:
result = service_module._get_message_context_by_conversation(
cast(sessionmaker[Session], session_maker),
message_session_maker,
conversation_id="conv-1",
workflow_run_id="run-1",
)
@ -163,22 +189,22 @@ class TestWorkflowEventSnapshotHelpers:
assert result is None
def test_get_message_context_by_conversation_should_default_created_at_to_zero_when_message_has_no_timestamp(
self,
self, message_session_maker: sessionmaker[Session]
) -> None:
message = SimpleNamespace(
id="msg-1",
conversation_id="conv-1",
created_at=None,
answer="answer",
)
session = SimpleNamespace(scalar=MagicMock(return_value=message))
session_maker = _SessionMaker(session)
_persist_message(message_session_maker)
result = service_module._get_message_context_by_conversation(
cast(sessionmaker[Session], session_maker),
conversation_id="conv-1",
workflow_run_id="run-1",
)
def clear_created_at(message: Message, _context: Any) -> None:
message.created_at = None # type: ignore[assignment]
event.listen(Message, "load", clear_created_at)
try:
result = service_module._get_message_context_by_conversation(
message_session_maker,
conversation_id="conv-1",
workflow_run_id="run-1",
)
finally:
event.remove(Message, "load", clear_created_at)
assert result is not None
assert result.created_at == 0

12
api/uv.lock generated
View File

@ -2710,14 +2710,14 @@ wheels = [
[[package]]
name = "gitpython"
version = "3.1.50"
version = "3.1.52"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "gitdb" },
]
sdist = { url = "https://files.pythonhosted.org/packages/33/f6/354ae6491228b5eb40e10d89c4d13c651fe1cf7556e35ebdded50cff57ce/gitpython-3.1.50.tar.gz", hash = "sha256:80da2d12504d52e1f998772dc5baf6e553f8d2fcfe1fcc226c9d9a2ee3372dcc", size = 219798, upload-time = "2026-05-06T04:01:26.571Z" }
sdist = { url = "https://files.pythonhosted.org/packages/e5/fd/df0bafa4eb5ea2f51e1adee9f7a94c8e62c5d180e65117045dfca3439c8a/gitpython-3.1.52.tar.gz", hash = "sha256:de0a8ad86274c6e75ae8b37dd055ba68f19818c813108642263227b20775b48e", size = 223726, upload-time = "2026-07-16T03:15:59.599Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/20/7a/1c6e3562dfd8950adbb11ffbc65d21e7c89d01a6e4f137fa981056de25c5/gitpython-3.1.50-py3-none-any.whl", hash = "sha256:d352abe2908d07355014abdd21ddf798c2a961469239afec4962e9da884858f9", size = 212507, upload-time = "2026-05-06T04:01:23.799Z" },
{ url = "https://files.pythonhosted.org/packages/8d/90/04dff7c1e176bb1c3011ef1647393d368790da710d8dde1cdcfad301f45a/gitpython-3.1.52-py3-none-any.whl", hash = "sha256:79a36ee1f83523214a3f72d56cf1c4e490d577dc61af77e43dfe5862bd9da01a", size = 215366, upload-time = "2026-07-16T03:15:58.239Z" },
]
[[package]]
@ -5119,11 +5119,11 @@ wheels = [
[[package]]
name = "pyasn1"
version = "0.6.3"
version = "0.6.4"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/5c/5f/6583902b6f79b399c9c40674ac384fd9cd77805f9e6205075f828ef11fb2/pyasn1-0.6.3.tar.gz", hash = "sha256:697a8ecd6d98891189184ca1fa05d1bb00e2f84b5977c481452050549c8a72cf", size = 148685, upload-time = "2026-03-17T01:06:53.382Z" }
sdist = { url = "https://files.pythonhosted.org/packages/a4/9a/23310166d960def5897e91fe20e5b724601b02a22e84ba1f94232c0b7f67/pyasn1-0.6.4.tar.gz", hash = "sha256:9c447d8431c947fe4c8febc4ed9e760bc29011a5b01e5c74b67025bd9fb8ce81", size = 151262, upload-time = "2026-07-09T01:12:33.988Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/5d/a0/7d793dce3fa811fe047d6ae2431c672364b462850c6235ae306c0efd025f/pyasn1-0.6.3-py3-none-any.whl", hash = "sha256:a80184d120f0864a52a073acc6fc642847d0be408e7c7252f31390c0f4eadcde", size = 83997, upload-time = "2026-03-17T01:06:52.036Z" },
{ url = "https://files.pythonhosted.org/packages/9a/3b/6163796d69c3977d1e4287bea4a6979161cbbdd170ebb430511e8e1999ce/pyasn1-0.6.4-py3-none-any.whl", hash = "sha256:deda9277cfd454080ec40b207fb6df82206a3a2688735233cdcd8d3d565f088b", size = 84410, upload-time = "2026-07-09T01:12:32.92Z" },
]
[[package]]

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