mirror of
https://github.com/langgenius/dify.git
synced 2026-09-08 11:04:27 +08:00
fix: fix conflict
This commit is contained in:
commit
9237f2a14a
@ -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
61
.github/CODEOWNERS
vendored
@ -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
|
||||
|
||||
2
.github/workflows/main-ci.yml
vendored
2
.github/workflows/main-ci.yml
vendored
@ -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:
|
||||
|
||||
16
.github/workflows/post-merge.yml
vendored
16
.github/workflows/post-merge.yml
vendored
@ -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/**'
|
||||
|
||||
11
.github/workflows/web-e2e.yml
vendored
11
.github/workflows/web-e2e.yml
vendored
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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()
|
||||
|
||||
|
||||
@ -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."
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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__:
|
||||
|
||||
@ -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"
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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)
|
||||
|
||||
|
||||
@ -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,
|
||||
)
|
||||
|
||||
@ -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.
|
||||
|
||||
@ -1,5 +1,5 @@
|
||||
{
|
||||
"commit": "4310e2d582d25e7de58183f27720afab01e123cf",
|
||||
"openapiSha256": "5827ca930ce38462bfd1b2bef387efbf37eb7ffcaedde4558af2fbaeccbfbc4b",
|
||||
"commit": "a0f50470612cc0b3656f89e4f2435aaf412e6e3b",
|
||||
"openapiSha256": "f18910e9c45a64f0855e0643a7a626fb2889021b4f943458de86c6bd2469facb",
|
||||
"repository": "https://github.com/langgenius/knowledge-fs"
|
||||
}
|
||||
|
||||
@ -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)
|
||||
|
||||
|
||||
|
||||
@ -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."
|
||||
|
||||
@ -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")
|
||||
@ -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)
|
||||
|
||||
|
||||
|
||||
@ -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}},
|
||||
},
|
||||
|
||||
@ -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 |
|
||||
|
||||
@ -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 |
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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]:
|
||||
|
||||
@ -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]
|
||||
|
||||
502
api/services/knowledge_fs_operations.py
Normal file
502
api/services/knowledge_fs_operations.py
Normal 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}"),
|
||||
)
|
||||
@ -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,
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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"]
|
||||
|
||||
@ -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() == []
|
||||
|
||||
@ -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,
|
||||
)
|
||||
|
||||
@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
@ -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"] == []
|
||||
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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"
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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"),
|
||||
@ -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()
|
||||
|
||||
@ -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):
|
||||
@ -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"
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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",
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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"):
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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"),
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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."""
|
||||
|
||||
@ -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)
|
||||
@ -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
|
||||
|
||||
67
api/tests/unit_tests/controllers/web/test_site.py
Normal file
67
api/tests/unit_tests/controllers/web/test_site.py
Normal 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()
|
||||
@ -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
|
||||
|
||||
@ -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])
|
||||
|
||||
@ -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"
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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"),
|
||||
[
|
||||
|
||||
@ -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"}},
|
||||
|
||||
@ -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",
|
||||
|
||||
@ -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)
|
||||
|
||||
|
||||
|
||||
@ -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"])
|
||||
}
|
||||
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@ -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,
|
||||
)
|
||||
@ -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()
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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]
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@ -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
|
||||
|
||||
@ -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
|
||||
@ -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"
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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
12
api/uv.lock
generated
@ -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
Loading…
Reference in New Issue
Block a user